Weiyun1025 commited on
Commit
201c300
·
verified ·
1 Parent(s): e4a4261

Upload folder using huggingface_hub

Browse files
Files changed (4) hide show
  1. config.json +63 -0
  2. dflash.py +617 -0
  3. dspark.py +385 -0
  4. model.safetensors +3 -0
config.json ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "DSparkDraftModel"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "auto_map": {
8
+ "AutoModel": "dspark.DSparkDraftModel"
9
+ },
10
+ "block_size": 7,
11
+ "bos_token_id": null,
12
+ "confidence_head_with_markov": true,
13
+ "dflash_config": {
14
+ "mask_token_id": 151675,
15
+ "target_layer_ids": [
16
+ 1,
17
+ 7,
18
+ 14,
19
+ 20,
20
+ 26,
21
+ 32,
22
+ 39,
23
+ 45
24
+ ],
25
+ "use_mask_embedding": true
26
+ },
27
+ "dtype": "bfloat16",
28
+ "enable_confidence_head": true,
29
+ "eos_token_id": 151645,
30
+ "head_dim": 128,
31
+ "hidden_act": "silu",
32
+ "hidden_size": 4096,
33
+ "initializer_range": 0.02,
34
+ "intermediate_size": 4096,
35
+ "layer_types": [
36
+ "sliding_attention",
37
+ "sliding_attention",
38
+ "sliding_attention",
39
+ "sliding_attention",
40
+ "sliding_attention"
41
+ ],
42
+ "markov_head_type": "vanilla",
43
+ "markov_rank": 256,
44
+ "max_position_embeddings": 1048576,
45
+ "max_window_layers": 28,
46
+ "model_type": "qwen3",
47
+ "num_attention_heads": 32,
48
+ "num_hidden_layers": 5,
49
+ "num_key_value_heads": 4,
50
+ "num_target_layers": 48,
51
+ "pad_token_id": 151643,
52
+ "rms_norm_eps": 1e-05,
53
+ "rope_parameters": {
54
+ "rope_theta": 10000,
55
+ "rope_type": "default"
56
+ },
57
+ "sliding_window": 1024,
58
+ "tie_word_embeddings": false,
59
+ "transformers_version": "5.12.1",
60
+ "use_cache": true,
61
+ "use_sliding_window": true,
62
+ "vocab_size": 152576
63
+ }
dflash.py ADDED
@@ -0,0 +1,617 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Callable, Optional
2
+
3
+ import torch
4
+ from torch import nn
5
+ from transformers import DynamicCache
6
+ from transformers.cache_utils import Cache
7
+ from transformers.modeling_outputs import CausalLMOutputWithPast
8
+ from transformers.models.qwen3.modeling_qwen3 import (
9
+ ALL_ATTENTION_FUNCTIONS,
10
+ FlashAttentionKwargs,
11
+ GradientCheckpointingLayer,
12
+ Qwen3Config,
13
+ Qwen3MLP,
14
+ Qwen3PreTrainedModel,
15
+ Qwen3RMSNorm,
16
+ Qwen3RotaryEmbedding,
17
+ eager_attention_forward,
18
+ rotate_half,
19
+ )
20
+ from typing_extensions import Tuple, Unpack
21
+
22
+
23
+ def sample(logits: torch.Tensor, temperature: float = 0.0) -> torch.Tensor:
24
+ if temperature < 1e-5:
25
+ return torch.argmax(logits, dim=-1)
26
+ bsz, seq_len, vocab_size = logits.shape
27
+ logits = logits.view(-1, vocab_size)
28
+ logits = logits / temperature
29
+ probs = torch.softmax(logits, dim=-1)
30
+ return torch.multinomial(probs, num_samples=1).view(bsz, seq_len)
31
+
32
+
33
+ def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):
34
+ cos = cos.unsqueeze(unsqueeze_dim)
35
+ sin = sin.unsqueeze(unsqueeze_dim)
36
+ rotary_dim = cos.size(-1)
37
+ if rotary_dim > q.size(-1) or rotary_dim > k.size(-1):
38
+ raise ValueError(
39
+ f"RoPE dim ({rotary_dim}) exceeds q/k dim ({q.size(-1)}, {k.size(-1)})."
40
+ )
41
+
42
+ q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]
43
+ k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]
44
+ q_len = q.size(-2)
45
+ q_rot = (q_rot * cos[..., -q_len:, :]) + (
46
+ rotate_half(q_rot) * sin[..., -q_len:, :]
47
+ )
48
+ k_rot = (k_rot * cos) + (rotate_half(k_rot) * sin)
49
+ q_embed = torch.cat((q_rot, q_pass), dim=-1)
50
+ k_embed = torch.cat((k_rot, k_pass), dim=-1)
51
+ return q_embed, k_embed
52
+
53
+
54
+ def apply_rotary_single(x, cos, sin, unsqueeze_dim=1):
55
+ """Apply (partial) RoPE to a single tensor whose seq length matches cos/sin.
56
+
57
+ Used by the ``use_target_kv`` path, where only the draft's own (query and
58
+ in-block noise-key) tokens need the draft RoPE — the target-provided context
59
+ K is already rotated in the target's space and must be left untouched.
60
+ ``x`` is ``[b, heads, L, head_dim]``; ``cos``/``sin`` are ``[b, L, rotary_dim]``.
61
+ """
62
+ cos = cos.unsqueeze(unsqueeze_dim)
63
+ sin = sin.unsqueeze(unsqueeze_dim)
64
+ rotary_dim = cos.size(-1)
65
+ if rotary_dim > x.size(-1):
66
+ raise ValueError(
67
+ f"RoPE dim ({rotary_dim}) exceeds tensor dim ({x.size(-1)})."
68
+ )
69
+ x_rot, x_pass = x[..., :rotary_dim], x[..., rotary_dim:]
70
+ x_rot = (x_rot * cos) + (rotate_half(x_rot) * sin)
71
+ return torch.cat((x_rot, x_pass), dim=-1)
72
+
73
+
74
+ class Qwen3DFlashAttention(nn.Module):
75
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
76
+
77
+ def __init__(self, config: Qwen3Config, layer_idx: int):
78
+ super().__init__()
79
+ self.config = config
80
+ self.layer_idx = layer_idx
81
+ self.head_dim = getattr(
82
+ config, "head_dim", config.hidden_size // config.num_attention_heads
83
+ )
84
+ self.num_key_value_groups = (
85
+ config.num_attention_heads // config.num_key_value_heads
86
+ )
87
+ self.scaling = self.head_dim**-0.5
88
+ self.attention_dropout = config.attention_dropout
89
+ self.is_causal = False
90
+ self.q_proj = nn.Linear(
91
+ config.hidden_size,
92
+ config.num_attention_heads * self.head_dim,
93
+ bias=config.attention_bias,
94
+ )
95
+ self.k_proj = nn.Linear(
96
+ config.hidden_size,
97
+ config.num_key_value_heads * self.head_dim,
98
+ bias=config.attention_bias,
99
+ )
100
+ self.v_proj = nn.Linear(
101
+ config.hidden_size,
102
+ config.num_key_value_heads * self.head_dim,
103
+ bias=config.attention_bias,
104
+ )
105
+ self.o_proj = nn.Linear(
106
+ config.num_attention_heads * self.head_dim,
107
+ config.hidden_size,
108
+ bias=config.attention_bias,
109
+ )
110
+ self.q_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)
111
+ self.k_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)
112
+ # target-KV consumption mode (see DFlashDraftModel): when target_kv is
113
+ # supplied, "inject" adds it as a residual on top of the draft's own
114
+ # k_proj/v_proj context K/V, whereas the default "replace" (use_target_kv)
115
+ # uses the target's KV directly. Read per-attention so forward can branch.
116
+ _dflash_cfg = getattr(config, "dflash_config", {}) or {}
117
+ self.use_target_kv_inject = bool(_dflash_cfg.get("use_target_kv_inject", False))
118
+ layer_types = getattr(config, "layer_types", None)
119
+ is_sliding_layer = (
120
+ isinstance(layer_types, (list, tuple))
121
+ and layer_idx < len(layer_types)
122
+ and layer_types[layer_idx] == "sliding_attention"
123
+ )
124
+ self.sliding_window = config.sliding_window if is_sliding_layer else None
125
+
126
+ def forward(
127
+ self,
128
+ hidden_states: torch.Tensor,
129
+ target_hidden: torch.Tensor,
130
+ position_embeddings: tuple[torch.Tensor, torch.Tensor],
131
+ attention_mask: Optional[torch.Tensor],
132
+ past_key_values: Optional[Cache] = None,
133
+ cache_position: Optional[torch.LongTensor] = None,
134
+ target_kv: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
135
+ **kwargs: Unpack[FlashAttentionKwargs],
136
+ ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
137
+ bsz, q_len = hidden_states.shape[:-1]
138
+ cos, sin = position_embeddings
139
+ q = self.q_proj(hidden_states)
140
+ q = q.view(bsz, q_len, -1, self.head_dim)
141
+ q = self.q_norm(q).transpose(1, 2)
142
+
143
+ if target_kv is not None and not self.use_target_kv_inject:
144
+ # ---- use_target_kv (REPLACE) ------------------------------------
145
+ # Context K/V come straight from the *target* model's own KV: the
146
+ # provided K is already k_norm'd + RoPE'd (in the target's space,
147
+ # i.e. exactly what sits in the target KV cache) and V is the raw
148
+ # projected value. Only the in-block draft (noise) tokens are keyed
149
+ # by the draft's own k_proj/v_proj here, since the target has no KV
150
+ # for not-yet-generated block tokens.
151
+ k_ctx, v_ctx = target_kv # each [bsz, ctx_len, num_kv_heads, head_dim]
152
+ k_ctx = k_ctx.transpose(1, 2) # [bsz, nkv, ctx_len, head_dim]
153
+ v_ctx = v_ctx.transpose(1, 2)
154
+ k_noise = self.k_proj(hidden_states).view(bsz, q_len, -1, self.head_dim)
155
+ k_noise = self.k_norm(k_noise).transpose(1, 2) # [bsz, nkv, q_len, hd]
156
+ v_noise = (
157
+ self.v_proj(hidden_states)
158
+ .view(bsz, q_len, -1, self.head_dim)
159
+ .transpose(1, 2)
160
+ )
161
+ # The draft (noise) tokens live at the last q_len position ids; the
162
+ # target context K is already rotated, so only rotate q and k_noise.
163
+ cos_draft, sin_draft = cos[:, -q_len:, :], sin[:, -q_len:, :]
164
+ q = apply_rotary_single(q, cos_draft, sin_draft)
165
+ k_noise = apply_rotary_single(k_noise, cos_draft, sin_draft)
166
+ k = torch.cat([k_ctx, k_noise], dim=2) # [bsz, nkv, ctx_len+q_len, hd]
167
+ v = torch.cat([v_ctx, v_noise], dim=2)
168
+ else:
169
+ # ---- baseline, and use_target_kv_inject (baseline + residual) ----
170
+ ctx_len = target_hidden.shape[1]
171
+ k_ctx = self.k_proj(target_hidden)
172
+ k_noise = self.k_proj(hidden_states)
173
+ v_ctx = self.v_proj(target_hidden)
174
+ v_noise = self.v_proj(hidden_states)
175
+ k = torch.cat([k_ctx, k_noise], dim=1).view(
176
+ bsz, ctx_len + q_len, -1, self.head_dim
177
+ )
178
+ v = torch.cat([v_ctx, v_noise], dim=1).view(
179
+ bsz, ctx_len + q_len, -1, self.head_dim
180
+ )
181
+ k = self.k_norm(k).transpose(1, 2)
182
+ v = v.transpose(1, 2)
183
+ q, k = apply_rotary_pos_emb(q, k, cos, sin)
184
+ if target_kv is not None:
185
+ # INJECT: add the target's own K/V into the context slice as a
186
+ # residual on top of the draft's projected+normed+roped context
187
+ # K/V. (k/v are [bsz, nkv, ctx_len+q_len, hd]; the context is the
188
+ # leading ctx_len keys. target K is already roped in the target
189
+ # space, matching the draft's roped context via inherited RoPE.)
190
+ k_inj, v_inj = target_kv # each [bsz, ctx_len, nkv, hd]
191
+ ctxL = k_inj.shape[1]
192
+ k[:, :, :ctxL, :] = k[:, :, :ctxL, :] + k_inj.transpose(1, 2)
193
+ v[:, :, :ctxL, :] = v[:, :, :ctxL, :] + v_inj.transpose(1, 2)
194
+ if past_key_values is not None:
195
+ cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
196
+ k, v = past_key_values.update(k, v, self.layer_idx, cache_kwargs)
197
+ attn_fn: Callable = eager_attention_forward
198
+ if self.config._attn_implementation != "eager":
199
+ attn_fn = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]
200
+ attn_output, attn_weights = attn_fn(
201
+ self,
202
+ q,
203
+ k,
204
+ v,
205
+ attention_mask,
206
+ dropout=0.0 if not self.training else self.attention_dropout,
207
+ scaling=self.scaling,
208
+ sliding_window=self.sliding_window,
209
+ **kwargs,
210
+ )
211
+ attn_output = attn_output.reshape(bsz, q_len, -1)
212
+ attn_output = self.o_proj(attn_output)
213
+ return attn_output, attn_weights
214
+
215
+
216
+ class Qwen3DFlashDecoderLayer(GradientCheckpointingLayer):
217
+ def __init__(self, config: Qwen3Config, layer_idx: int):
218
+ super().__init__()
219
+ self.hidden_size = config.hidden_size
220
+ self.self_attn = Qwen3DFlashAttention(config=config, layer_idx=layer_idx)
221
+ self.mlp = Qwen3MLP(config)
222
+ self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
223
+ self.post_attention_layernorm = Qwen3RMSNorm(
224
+ config.hidden_size, eps=config.rms_norm_eps
225
+ )
226
+
227
+ def forward(
228
+ self,
229
+ target_hidden: Optional[torch.Tensor] = None,
230
+ hidden_states: Optional[torch.Tensor] = None,
231
+ attention_mask: Optional[torch.Tensor] = None,
232
+ position_ids: Optional[torch.LongTensor] = None,
233
+ past_key_value: Optional[Cache] = None,
234
+ output_attentions: Optional[bool] = False,
235
+ use_cache: Optional[bool] = False,
236
+ cache_position: Optional[torch.LongTensor] = None,
237
+ position_embeddings: Optional[
238
+ Tuple[torch.Tensor, torch.Tensor]
239
+ ] = None, # necessary, but kept here for BC
240
+ target_kv: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
241
+ **kwargs: Unpack[FlashAttentionKwargs],
242
+ ) -> Tuple[
243
+ torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]
244
+ ]:
245
+ residual = hidden_states
246
+ hidden_states = self.input_layernorm(hidden_states)
247
+ hidden_states = self.self_attn(
248
+ hidden_states=hidden_states,
249
+ target_hidden=target_hidden,
250
+ attention_mask=attention_mask,
251
+ position_ids=position_ids,
252
+ past_key_values=past_key_value,
253
+ output_attentions=output_attentions,
254
+ use_cache=use_cache,
255
+ cache_position=cache_position,
256
+ position_embeddings=position_embeddings,
257
+ target_kv=target_kv,
258
+ **kwargs,
259
+ )[0]
260
+ hidden_states = residual + hidden_states
261
+ residual = hidden_states
262
+ hidden_states = self.post_attention_layernorm(hidden_states)
263
+ hidden_states = self.mlp(hidden_states)
264
+ hidden_states = residual + hidden_states
265
+ return hidden_states
266
+
267
+
268
+ def build_target_layer_ids(num_target_layers: int, num_draft_layers: int):
269
+ if num_draft_layers == 1:
270
+ return [(num_target_layers // 2)]
271
+ start = 1
272
+ end = num_target_layers - 3
273
+ span = end - start
274
+ target_layer_ids = [
275
+ int(round(start + (i * span) / (num_draft_layers - 1)))
276
+ for i in range(num_draft_layers)
277
+ ]
278
+ return target_layer_ids
279
+
280
+
281
+ def extract_context_feature(
282
+ hidden_states: list[torch.Tensor],
283
+ layer_ids: Optional[list[int]],
284
+ ) -> torch.Tensor:
285
+ offset = 1
286
+ selected_states = []
287
+ for layer_id in layer_ids:
288
+ selected_states.append(hidden_states[layer_id + offset])
289
+ target_hidden = torch.cat(selected_states, dim=-1)
290
+ return target_hidden
291
+
292
+
293
+ class DFlashDraftModel(Qwen3PreTrainedModel):
294
+ config_class = Qwen3Config
295
+ _no_split_modules = ["Qwen3DFlashDecoderLayer"]
296
+
297
+ def __init__(self, config) -> None:
298
+ super().__init__(config)
299
+ self.config = config
300
+ self.layers = nn.ModuleList(
301
+ [
302
+ Qwen3DFlashDecoderLayer(config, layer_idx)
303
+ for layer_idx in range(config.num_hidden_layers)
304
+ ]
305
+ )
306
+ dflash_config = getattr(config, "dflash_config", {}) or {}
307
+ self.target_layer_ids = dflash_config.get(
308
+ "target_layer_ids",
309
+ build_target_layer_ids(config.num_target_layers, config.num_hidden_layers),
310
+ )
311
+ self.norm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
312
+ self.rotary_emb = Qwen3RotaryEmbedding(config)
313
+ # When use_target_kv is on, the draft's context K/V come directly from the
314
+ # target model's own per-layer KV (draft layer i ← target layer
315
+ # target_layer_ids[i]), so the fc/hidden_norm fusion of the aux hidden
316
+ # concat is not used and is not built. Draft layer i attends into the KV
317
+ # of target_layer_ids[i]; the list length already matches num draft layers.
318
+ self.use_target_kv = bool(dflash_config.get("use_target_kv", False))
319
+ # Alternative target-KV mode: instead of REPLACING the draft's context K/V
320
+ # with the target's KV, keep the draft's k_proj/v_proj(fc(target_hidden))
321
+ # and ADD the target KV as a residual injection. Keeps fc/hidden_norm.
322
+ self.use_target_kv_inject = bool(
323
+ dflash_config.get("use_target_kv_inject", False)
324
+ )
325
+ # Third mode: fuse ALL captured target layers' K/V (like the baseline fuses
326
+ # the aux HIDDEN layers) into one [B,S,H] feature via fc, then let each draft
327
+ # layer's own k_proj/v_proj project it — i.e. the target KV cache replaces
328
+ # the aux hidden as the fc input. Keeps the baseline attention path (each
329
+ # layer sees target_kv=None; context K/V = k_proj/v_proj of the fused KV).
330
+ self.use_target_kv_fuse = bool(dflash_config.get("use_target_kv_fuse", False))
331
+ if (
332
+ sum(
333
+ [
334
+ self.use_target_kv,
335
+ self.use_target_kv_inject,
336
+ self.use_target_kv_fuse,
337
+ ]
338
+ )
339
+ > 1
340
+ ):
341
+ raise ValueError(
342
+ "use_target_kv (replace) / use_target_kv_inject (residual add) / "
343
+ "use_target_kv_fuse (fuse KV as fc input) are mutually exclusive; "
344
+ "enable at most one."
345
+ )
346
+ if self.use_target_kv:
347
+ self.fc = None
348
+ self.hidden_norm = None
349
+ elif self.use_target_kv_fuse:
350
+ # fc input = concat over captured layers of flattened (K, V):
351
+ # len(target_layer_ids) * 2 * num_kv_heads * head_dim.
352
+ _hd = getattr(
353
+ config, "head_dim", config.hidden_size // config.num_attention_heads
354
+ )
355
+ _kv_feat = 2 * config.num_key_value_heads * _hd
356
+ self.fc = nn.Linear(
357
+ len(self.target_layer_ids) * _kv_feat,
358
+ config.hidden_size,
359
+ bias=False,
360
+ )
361
+ self.hidden_norm = Qwen3RMSNorm(
362
+ config.hidden_size, eps=config.rms_norm_eps
363
+ )
364
+ else:
365
+ self.fc = nn.Linear(
366
+ len(self.target_layer_ids) * config.hidden_size,
367
+ config.hidden_size,
368
+ bias=False,
369
+ )
370
+ self.hidden_norm = Qwen3RMSNorm(
371
+ config.hidden_size, eps=config.rms_norm_eps
372
+ )
373
+ self.block_size = config.block_size
374
+ self.mask_token_id = dflash_config.get("mask_token_id", None)
375
+ # Optional MiMo-style learned mask embedding. When enabled, the masked
376
+ # (to-be-predicted) block positions use this trained vector instead of
377
+ # the frozen target embed_tokens(mask_token_id). It is a normal draft
378
+ # parameter (optimized by the draft optimizer and saved into the draft
379
+ # checkpoint), and is additionally exported as `mask_embedding.pt` for
380
+ # the sglang DFlash worker, which injects it at each block's masked
381
+ # slots (noise_embedding[:, 1:, :]). Toggle via
382
+ # dflash_config["use_mask_embedding"] (CLI: --use-mask-embedding).
383
+ self.use_mask_embedding = bool(dflash_config.get("use_mask_embedding", False))
384
+ if self.use_mask_embedding:
385
+ self.mask_embedding = nn.Parameter(torch.zeros(config.hidden_size))
386
+ else:
387
+ self.register_parameter("mask_embedding", None)
388
+ self._offload_fc_input_enabled = False
389
+ self.post_init()
390
+
391
+ def set_offload_fc_input_enabled(self, enabled: bool) -> None:
392
+ self._offload_fc_input_enabled = enabled
393
+
394
+ def _pack_saved_tensor_to_cpu(self, tensor: torch.Tensor):
395
+ if tensor.device.type != "cuda":
396
+ return tensor, tensor.device
397
+ return tensor.to("cpu", non_blocking=True), tensor.device
398
+
399
+ @staticmethod
400
+ def _unpack_saved_tensor_from_cpu(packed):
401
+ cpu_tensor, device = packed
402
+ return cpu_tensor.to(device, non_blocking=True)
403
+
404
+ def forward(
405
+ self,
406
+ position_ids: torch.LongTensor,
407
+ attention_mask: Optional[torch.Tensor] = None,
408
+ noise_embedding: Optional[torch.Tensor] = None,
409
+ target_hidden: Optional[torch.Tensor] = None,
410
+ target_kv: Optional[list] = None,
411
+ past_key_values: Optional[Cache] = None,
412
+ use_cache: bool = False,
413
+ **kwargs,
414
+ ) -> CausalLMOutputWithPast:
415
+ hidden_states = noise_embedding
416
+
417
+ if self.mask_embedding is not None:
418
+ # Overwrite each block's masked slots (index % block_size != 0) with
419
+ # the learned mask embedding; each block's first slot keeps the real
420
+ # anchor-token embedding. Mirrors the sglang worker, which sets
421
+ # noise_embedding[:, 1:, :] = mask_embedding per block. Works for both
422
+ # training (length = n * block_size) and spec_generate (length =
423
+ # block_size), since the anchor always sits at index % block_size == 0.
424
+ seq_len = hidden_states.shape[1]
425
+ pos = torch.arange(seq_len, device=hidden_states.device)
426
+ is_mask_pos = (pos % self.block_size) != 0
427
+ hidden_states = torch.where(
428
+ is_mask_pos.view(1, seq_len, 1),
429
+ self.mask_embedding.to(hidden_states.dtype).view(1, 1, -1),
430
+ hidden_states,
431
+ )
432
+
433
+ needs_target_kv = (
434
+ self.use_target_kv or self.use_target_kv_inject or self.use_target_kv_fuse
435
+ )
436
+ if needs_target_kv:
437
+ # All three target-KV modes require the per-draft-layer target K/V.
438
+ if target_kv is None:
439
+ raise ValueError(
440
+ "use_target_kv / use_target_kv_inject / use_target_kv_fuse is "
441
+ "enabled but no target_kv was provided to "
442
+ "DFlashDraftModel.forward."
443
+ )
444
+ if len(target_kv) != len(self.layers):
445
+ raise ValueError(
446
+ f"target_kv has {len(target_kv)} entries but the draft has "
447
+ f"{len(self.layers)} layers; expected one (k, v) pair per layer."
448
+ )
449
+
450
+ if self.use_target_kv:
451
+ # REPLACE mode: the aux-hidden fusion (fc/hidden_norm) is not used.
452
+ pass
453
+ elif self.use_target_kv_fuse:
454
+ # FUSE mode: build the context feature from the captured target K/V
455
+ # instead of the aux hidden — concat every layer's flattened (K, V),
456
+ # fc-fuse to [B,S,H], then the baseline per-layer k_proj/v_proj project
457
+ # it (attention runs the baseline path, target_kv=None). NOTE the fused
458
+ # K carries the target's RoPE and the baseline re-applies the draft RoPE
459
+ # to k_ctx (double rotation on the K component); it is a fixed function
460
+ # of position that fc/k_proj learn around, but is a known subtlety.
461
+ kv_feats = []
462
+ for k_l, v_l in target_kv: # each [B, S, num_kv_heads, head_dim]
463
+ b, s = k_l.shape[:2]
464
+ kv_feats.append(k_l.reshape(b, s, -1))
465
+ kv_feats.append(v_l.reshape(b, s, -1))
466
+ kv_concat = torch.cat(kv_feats, dim=-1) # [B, S, K*2*nkv*hd]
467
+ target_hidden = self.hidden_norm(self.fc(kv_concat))
468
+ elif self._offload_fc_input_enabled:
469
+ # Offload tensors saved by fc/norm autograd so later attention layers
470
+ # can reuse the GPU memory; they are copied back during backward.
471
+ with torch.autograd.graph.saved_tensors_hooks(
472
+ self._pack_saved_tensor_to_cpu,
473
+ self._unpack_saved_tensor_from_cpu,
474
+ ):
475
+ target_hidden = self.hidden_norm(self.fc(target_hidden))
476
+ else:
477
+ # baseline AND inject: fc-fuse the aux hidden into the context feature.
478
+ target_hidden = self.hidden_norm(self.fc(target_hidden))
479
+
480
+ position_embeddings = self.rotary_emb(hidden_states, position_ids)
481
+ # target_kv is consumed by the attention (replace/inject) only; in fuse mode
482
+ # it has already been folded into target_hidden above, so the attention runs
483
+ # the plain baseline path.
484
+ attn_uses_target_kv = self.use_target_kv or self.use_target_kv_inject
485
+ for layer_idx, layer in enumerate(self.layers):
486
+ layer_attention_mask = attention_mask
487
+ if isinstance(attention_mask, dict):
488
+ layer_attention_mask = (
489
+ attention_mask["sliding"]
490
+ if layer.self_attn.sliding_window is not None
491
+ else attention_mask["full"]
492
+ )
493
+ hidden_states = layer(
494
+ hidden_states=hidden_states,
495
+ target_hidden=None if self.use_target_kv else target_hidden,
496
+ attention_mask=layer_attention_mask,
497
+ position_ids=position_ids,
498
+ past_key_value=past_key_values,
499
+ use_cache=use_cache,
500
+ position_embeddings=position_embeddings,
501
+ target_kv=target_kv[layer_idx] if attn_uses_target_kv else None,
502
+ **kwargs,
503
+ )
504
+ return self.norm(hidden_states)
505
+
506
+ @torch.inference_mode()
507
+ def spec_generate(
508
+ self,
509
+ target: nn.Module,
510
+ input_ids: torch.LongTensor,
511
+ max_new_tokens: int,
512
+ stop_token_ids: list[int],
513
+ temperature: float,
514
+ ):
515
+ self.eval()
516
+ num_input_tokens = input_ids.shape[1]
517
+ max_length = num_input_tokens + max_new_tokens
518
+
519
+ block_size = self.block_size
520
+ output_ids = torch.full(
521
+ (1, max_length + block_size),
522
+ self.mask_token_id,
523
+ dtype=torch.long,
524
+ device=target.device,
525
+ )
526
+ position_ids = torch.arange(
527
+ output_ids.shape[1], device=target.device
528
+ ).unsqueeze(0)
529
+
530
+ past_key_values_target = DynamicCache()
531
+ past_key_values_draft = DynamicCache()
532
+
533
+ # Prefill stage
534
+ output = target(
535
+ input_ids,
536
+ position_ids=position_ids[:, :num_input_tokens],
537
+ past_key_values=past_key_values_target,
538
+ use_cache=True,
539
+ logits_to_keep=1,
540
+ output_hidden_states=True,
541
+ )
542
+
543
+ output_ids[:, :num_input_tokens] = input_ids
544
+ output_ids[:, num_input_tokens : num_input_tokens + 1] = sample(
545
+ output.logits, temperature
546
+ )
547
+ target_hidden = extract_context_feature(
548
+ output.hidden_states, self.target_layer_ids
549
+ )
550
+
551
+ # Decode stage
552
+ acceptance_lengths = []
553
+ start = input_ids.shape[1]
554
+ while start < max_length:
555
+ block_output_ids = output_ids[:, start : start + block_size].clone()
556
+ block_position_ids = position_ids[:, start : start + block_size]
557
+ noise_embedding = target.model.embed_tokens(block_output_ids)
558
+ draft_logits = target.lm_head(
559
+ self(
560
+ target_hidden=target_hidden,
561
+ noise_embedding=noise_embedding,
562
+ position_ids=position_ids[
563
+ :, past_key_values_draft.get_seq_length() : start + block_size
564
+ ],
565
+ past_key_values=past_key_values_draft,
566
+ use_cache=True,
567
+ is_causal=False,
568
+ )[:, -block_size + 1 :, :]
569
+ )
570
+ past_key_values_draft.crop(start)
571
+ block_output_ids[:, 1:] = sample(draft_logits)
572
+
573
+ output = target(
574
+ block_output_ids,
575
+ position_ids=block_position_ids,
576
+ past_key_values=past_key_values_target,
577
+ use_cache=True,
578
+ output_hidden_states=True,
579
+ )
580
+
581
+ posterior = sample(output.logits, temperature)
582
+ acceptance_length = (
583
+ (block_output_ids[:, 1:] == posterior[:, :-1])
584
+ .cumprod(dim=1)
585
+ .sum(dim=1)[0]
586
+ .item()
587
+ )
588
+ output_ids[:, start : start + acceptance_length + 1] = block_output_ids[
589
+ :, : acceptance_length + 1
590
+ ]
591
+ output_ids[:, start + acceptance_length + 1] = posterior[
592
+ :, acceptance_length
593
+ ]
594
+ start += acceptance_length + 1
595
+ past_key_values_target.crop(start)
596
+ target_hidden = extract_context_feature(
597
+ output.hidden_states, self.target_layer_ids
598
+ )[:, : acceptance_length + 1, :]
599
+ acceptance_lengths.append(acceptance_length + 1)
600
+ if stop_token_ids is not None and any(
601
+ stop_token_id in output_ids[:, num_input_tokens:]
602
+ for stop_token_id in stop_token_ids
603
+ ):
604
+ break
605
+ output_ids = output_ids[:, :max_length]
606
+ output_ids = output_ids[:, output_ids[0] != self.mask_token_id]
607
+ if stop_token_ids is not None:
608
+ stop_token_ids = torch.tensor(stop_token_ids, device=output_ids.device)
609
+ stop_token_indices = torch.isin(
610
+ output_ids[0][num_input_tokens:], stop_token_ids
611
+ ).nonzero(as_tuple=True)[0]
612
+ if stop_token_indices.numel() > 0:
613
+ output_ids = output_ids[
614
+ :, : num_input_tokens + stop_token_indices[0] + 1
615
+ ]
616
+
617
+ return output_ids
dspark.py ADDED
@@ -0,0 +1,385 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ """DSpark draft model: DFlash backbone + EAGLE-style Markov and confidence heads.
3
+
4
+ DSpark shares SpecForge's DFlash block-diffusion drafter (dual-source KV
5
+ injection via :class:`DFlashDraftModel`, anchor sampling, MASK-token noise
6
+ stream) and adds two heads on top:
7
+
8
+ - Markov head: a learned low-rank bias added to the draft logits, conditioned
9
+ on the (teacher-forced) previous token. Three variants are supported, exactly
10
+ mirroring DeepSpec: ``vanilla`` (memoryless bigram), ``gated`` (token-gated),
11
+ and ``rnn`` (recurrent state across within-block positions).
12
+ - Confidence head (AcceptRatePredictor): predicts a per-draft-position
13
+ acceptance probability, trained against the empirical draft-vs-target accept
14
+ rate (used at inference for adaptive block length).
15
+
16
+ The Markov / confidence / accept-rate modeling is ported to match DeepSeek's
17
+ DeepSpec one-for-one (``deepspec/modeling/dspark/{markov_head,common}.py``, MIT
18
+ License). SpecForge structural differences (load-bearing):
19
+ - There is no ``DFlashConfig``; SpecForge's :class:`DFlashDraftModel` uses a
20
+ plain ``Qwen3Config`` plus a ``config.dflash_config`` dict. So
21
+ :class:`DSparkConfig` subclasses ``Qwen3Config`` and declares the DSpark
22
+ fields as top-level attributes; DFlash-carried fields (``block_size``,
23
+ ``num_target_layers``, ``dflash_config``) stay as before.
24
+ - The draft model has no ``embed_tokens`` / ``lm_head`` of its own (they live on
25
+ the target and are passed into the online wrapper). The heads only depend on
26
+ ``config.hidden_size`` / ``config.vocab_size``, so this does not matter for
27
+ construction.
28
+ - DeepSpec builds the heads *before* ``post_init`` so the HF initializer
29
+ (normal, std=initializer_range) covers them. SpecForge's base ``__init__``
30
+ runs ``post_init`` before the DSpark heads exist, so we re-apply
31
+ ``_init_weights`` to the heads here to reproduce DeepSpec's initialization
32
+ exactly (without this, ``markov_w1`` would keep the nn.Embedding default
33
+ N(0,1) and the Markov bias would be huge at init).
34
+ """
35
+
36
+ from typing import Optional
37
+
38
+ import torch
39
+ import torch.nn as nn
40
+ from transformers.models.qwen3.modeling_qwen3 import Qwen3Config
41
+
42
+ from specforge.modeling.draft.dflash import DFlashDraftModel
43
+
44
+
45
+ def _sample_tokens(logits: torch.Tensor, temperature: float = 0.0) -> torch.Tensor:
46
+ """Greedy (temperature<1e-5) or multinomial sampling over the last dim.
47
+
48
+ Mirrors DeepSpec ``deepspec/utils/sampling.py::sample_tokens``. Used only by
49
+ the heads' inference-time ``sample_block_tokens`` (not the training forward).
50
+ """
51
+ if temperature < 1e-5:
52
+ return torch.argmax(logits, dim=-1)
53
+ bsz, seq_len, vocab_size = logits.shape
54
+ flat_logits = logits.reshape(-1, vocab_size) / temperature
55
+ probs = torch.softmax(flat_logits, dim=-1)
56
+ return torch.multinomial(probs, num_samples=1).reshape(bsz, seq_len)
57
+
58
+
59
+ class DSparkConfig(Qwen3Config):
60
+ """Configuration for the DSpark draft model.
61
+
62
+ Extends ``Qwen3Config``. DSpark-specific fields are declared here; the
63
+ DFlash-carried fields (``block_size``, ``num_target_layers``, and the nested
64
+ ``dflash_config`` dict holding ``target_layer_ids`` / ``mask_token_id``) are
65
+ consumed by the :class:`DFlashDraftModel` base ``__init__`` and must be
66
+ present on the config object before constructing the model.
67
+ """
68
+
69
+ model_type = "dspark"
70
+
71
+ def __init__(
72
+ self,
73
+ markov_rank: int = 256,
74
+ markov_head_type: str = "vanilla",
75
+ enable_confidence_head: bool = True,
76
+ confidence_head_with_markov: bool = True,
77
+ **kwargs,
78
+ ):
79
+ super().__init__(**kwargs)
80
+ self.markov_rank = markov_rank
81
+ self.markov_head_type = markov_head_type
82
+ self.enable_confidence_head = enable_confidence_head
83
+ self.confidence_head_with_markov = confidence_head_with_markov
84
+
85
+
86
+ class VanillaMarkov(nn.Module):
87
+ """Memoryless low-rank learned bigram bias added to the draft logits.
88
+
89
+ Ported from DeepSpec ``deepspec/modeling/dspark/markov_head.py``.
90
+ """
91
+
92
+ def __init__(self, *, vocab_size: int, markov_rank: int):
93
+ super().__init__()
94
+ self.vocab_size = int(vocab_size)
95
+ self.markov_rank = int(markov_rank)
96
+ self.markov_head_type = "vanilla"
97
+ assert (
98
+ self.markov_rank > 0
99
+ ), f"VanillaMarkov requires markov_rank > 0, got {self.markov_rank}."
100
+ self.markov_w1 = nn.Embedding(self.vocab_size, self.markov_rank)
101
+ self.markov_w2 = nn.Linear(self.markov_rank, self.vocab_size, bias=False)
102
+
103
+ def get_prev_embeddings(self, token_ids: torch.Tensor) -> torch.Tensor:
104
+ return self.markov_w1(token_ids.long())
105
+
106
+ def project_bias(self, latent_states: torch.Tensor) -> torch.Tensor:
107
+ return self.markov_w2(latent_states)
108
+
109
+ def compute_step_bias(
110
+ self,
111
+ token_ids: torch.Tensor,
112
+ hidden_states: Optional[torch.Tensor] = None,
113
+ ) -> torch.Tensor:
114
+ del hidden_states
115
+ return self.project_bias(self.get_prev_embeddings(token_ids))
116
+
117
+ def apply_step_logits(
118
+ self,
119
+ logits: torch.Tensor,
120
+ *,
121
+ token_ids: torch.Tensor,
122
+ hidden_states: Optional[torch.Tensor] = None,
123
+ ) -> torch.Tensor:
124
+ return logits + self.compute_step_bias(token_ids, hidden_states)
125
+
126
+ def apply_block_logits(
127
+ self,
128
+ base_logits: torch.Tensor,
129
+ *,
130
+ token_ids: torch.Tensor,
131
+ hidden_states: Optional[torch.Tensor] = None,
132
+ ) -> torch.Tensor:
133
+ if base_logits.size(2) == 0:
134
+ return base_logits
135
+ return base_logits + self.compute_step_bias(token_ids, hidden_states)
136
+
137
+ def sample_block_tokens(
138
+ self,
139
+ base_logits: torch.Tensor,
140
+ *,
141
+ first_prev_token_ids: torch.Tensor,
142
+ hidden_states: Optional[torch.Tensor] = None,
143
+ temperature: float = 0.0,
144
+ ):
145
+ batch_size, proposal_len = base_logits.shape[:2]
146
+ if proposal_len == 0:
147
+ empty_tokens = torch.empty(
148
+ batch_size, 0, dtype=torch.long, device=base_logits.device
149
+ )
150
+ return empty_tokens, base_logits
151
+
152
+ sampled_tokens = []
153
+ corrected_logits = []
154
+ prev_token_ids = first_prev_token_ids.long()
155
+ for step_idx in range(proposal_len):
156
+ step_hidden = (
157
+ None if hidden_states is None else hidden_states[:, step_idx, ...]
158
+ )
159
+ step_logits = self.apply_step_logits(
160
+ base_logits[:, step_idx, :],
161
+ token_ids=prev_token_ids,
162
+ hidden_states=step_hidden,
163
+ )
164
+ corrected_logits.append(step_logits.unsqueeze(1))
165
+ next_token_ids = _sample_tokens(
166
+ step_logits.unsqueeze(1), temperature=temperature
167
+ ).squeeze(1)
168
+ sampled_tokens.append(next_token_ids)
169
+ prev_token_ids = next_token_ids
170
+ return torch.stack(sampled_tokens, dim=1), torch.cat(corrected_logits, dim=1)
171
+
172
+
173
+ class GatedMarkovHead(VanillaMarkov):
174
+ """Token-gated Markov head (DeepSpec ``gated``).
175
+
176
+ The previous-token embedding is gated by a sigmoid of [hidden; prev_emb]
177
+ before projection, letting the backbone hidden modulate the bigram bias.
178
+ """
179
+
180
+ def __init__(self, *, vocab_size: int, markov_rank: int, hidden_size: int):
181
+ super().__init__(vocab_size=vocab_size, markov_rank=markov_rank)
182
+ self.markov_head_type = "gated"
183
+ self.gate_proj = nn.Linear(hidden_size + markov_rank, markov_rank)
184
+
185
+ def compute_gate(
186
+ self,
187
+ token_ids: torch.Tensor,
188
+ hidden_states: Optional[torch.Tensor],
189
+ ) -> torch.Tensor:
190
+ assert hidden_states is not None
191
+ prev_embeddings = self.get_prev_embeddings(token_ids)
192
+ gate_inputs = torch.cat([hidden_states, prev_embeddings], dim=-1)
193
+ return torch.sigmoid(self.gate_proj(gate_inputs))
194
+
195
+ def compute_step_bias(
196
+ self,
197
+ token_ids: torch.Tensor,
198
+ hidden_states: Optional[torch.Tensor] = None,
199
+ ) -> torch.Tensor:
200
+ prev_embeddings = self.get_prev_embeddings(token_ids)
201
+ gate = self.compute_gate(token_ids, hidden_states).to(
202
+ dtype=prev_embeddings.dtype
203
+ )
204
+ return self.project_bias(gate * prev_embeddings)
205
+
206
+
207
+ class RNNHead(VanillaMarkov):
208
+ """Recurrent Markov head (DeepSpec ``rnn``).
209
+
210
+ Maintains a GRU-like recurrent state across within-block positions, so
211
+ position k can access the full prefix history x_{<k}.
212
+ """
213
+
214
+ def __init__(self, *, vocab_size: int, markov_rank: int, hidden_size: int):
215
+ super().__init__(vocab_size=vocab_size, markov_rank=markov_rank)
216
+ self.markov_head_type = "rnn"
217
+ self.hidden_size = hidden_size
218
+ self.state_size = markov_rank
219
+ # [s_{k-1}; W1[x_{k-1}]; h_k] -> [gate; candidate; output]
220
+ self.joint_proj = nn.Linear(2 * markov_rank + hidden_size, 3 * markov_rank)
221
+
222
+ def _rnn_step(
223
+ self,
224
+ state: torch.Tensor,
225
+ prev_embeddings: torch.Tensor,
226
+ hidden_states: torch.Tensor,
227
+ ):
228
+ z = torch.cat([state, prev_embeddings, hidden_states], dim=-1)
229
+ proj = self.joint_proj(z)
230
+ gate_raw, candidate_raw, output_raw = proj.chunk(3, dim=-1)
231
+ gate = torch.sigmoid(gate_raw)
232
+ candidate = torch.tanh(candidate_raw)
233
+ new_state = gate * state + (1.0 - gate) * candidate
234
+ bias = self.project_bias(torch.tanh(output_raw))
235
+ return new_state, bias
236
+
237
+ def compute_step_bias(
238
+ self,
239
+ token_ids: torch.Tensor,
240
+ hidden_states: Optional[torch.Tensor] = None,
241
+ ) -> torch.Tensor:
242
+ """Stateless single-step bias (state initialized to zero)."""
243
+ assert hidden_states is not None
244
+ prev_embeddings = self.get_prev_embeddings(token_ids)
245
+ state = torch.zeros_like(prev_embeddings)
246
+ _, bias = self._rnn_step(state, prev_embeddings, hidden_states)
247
+ return bias
248
+
249
+ def apply_block_logits(
250
+ self,
251
+ base_logits: torch.Tensor,
252
+ *,
253
+ token_ids: torch.Tensor,
254
+ hidden_states: Optional[torch.Tensor] = None,
255
+ ) -> torch.Tensor:
256
+ assert hidden_states is not None
257
+ block_size = base_logits.size(-2)
258
+ if block_size == 0:
259
+ return base_logits
260
+ leading_shape = base_logits.shape[:-2]
261
+ state = torch.zeros(
262
+ *leading_shape,
263
+ self.markov_rank,
264
+ device=base_logits.device,
265
+ dtype=hidden_states.dtype,
266
+ )
267
+ output_logits = []
268
+ for k in range(block_size):
269
+ prev_emb = self.get_prev_embeddings(token_ids[..., k])
270
+ h_k = hidden_states[..., k, :]
271
+ state, bias = self._rnn_step(state, prev_emb, h_k)
272
+ output_logits.append(base_logits[..., k, :] + bias)
273
+ return torch.stack(output_logits, dim=-2)
274
+
275
+ def sample_block_tokens(
276
+ self,
277
+ base_logits: torch.Tensor,
278
+ *,
279
+ first_prev_token_ids: torch.Tensor,
280
+ hidden_states: Optional[torch.Tensor] = None,
281
+ temperature: float = 0.0,
282
+ ):
283
+ assert hidden_states is not None
284
+ batch_size, proposal_len = base_logits.shape[:2]
285
+ if proposal_len == 0:
286
+ empty_tokens = torch.empty(
287
+ batch_size, 0, dtype=torch.long, device=base_logits.device
288
+ )
289
+ return empty_tokens, base_logits
290
+ state = torch.zeros(
291
+ batch_size,
292
+ self.markov_rank,
293
+ device=base_logits.device,
294
+ dtype=hidden_states.dtype,
295
+ )
296
+ sampled_tokens = []
297
+ corrected_logits = []
298
+ prev_token_ids = first_prev_token_ids.long()
299
+ for step_idx in range(proposal_len):
300
+ prev_emb = self.get_prev_embeddings(prev_token_ids)
301
+ h_k = hidden_states[:, step_idx, :]
302
+ state, bias = self._rnn_step(state, prev_emb, h_k)
303
+ step_logits = base_logits[:, step_idx, :] + bias
304
+ corrected_logits.append(step_logits.unsqueeze(1))
305
+ next_token_ids = _sample_tokens(
306
+ step_logits.unsqueeze(1), temperature=temperature
307
+ ).squeeze(1)
308
+ sampled_tokens.append(next_token_ids)
309
+ prev_token_ids = next_token_ids
310
+ return torch.stack(sampled_tokens, dim=1), torch.cat(corrected_logits, dim=1)
311
+
312
+
313
+ class AcceptRatePredictor(nn.Module):
314
+ """Per-position acceptance-probability predictor (a single linear head).
315
+
316
+ Ported from DeepSpec ``deepspec/modeling/dspark/common.py``.
317
+ """
318
+
319
+ def __init__(self, input_dim: int):
320
+ super().__init__()
321
+ self.proj = nn.Linear(int(input_dim), 1)
322
+
323
+ def forward(self, features: torch.Tensor) -> torch.Tensor:
324
+ return self.proj(features).squeeze(-1)
325
+
326
+
327
+ def build_markov_head(config) -> Optional[nn.Module]:
328
+ markov_rank = int(getattr(config, "markov_rank", 0))
329
+ assert markov_rank >= 0, f"markov_rank must be >= 0, got {markov_rank}"
330
+ if markov_rank == 0:
331
+ return None
332
+
333
+ markov_head_type = str(getattr(config, "markov_head_type", "vanilla")).lower()
334
+ if markov_head_type == "vanilla":
335
+ return VanillaMarkov(vocab_size=config.vocab_size, markov_rank=markov_rank)
336
+ if markov_head_type == "gated":
337
+ return GatedMarkovHead(
338
+ vocab_size=config.vocab_size,
339
+ markov_rank=markov_rank,
340
+ hidden_size=config.hidden_size,
341
+ )
342
+ if markov_head_type == "rnn":
343
+ return RNNHead(
344
+ vocab_size=config.vocab_size,
345
+ markov_rank=markov_rank,
346
+ hidden_size=config.hidden_size,
347
+ )
348
+ raise AssertionError(f"Unsupported markov_head_type: {markov_head_type!r}")
349
+
350
+
351
+ class DSparkDraftModel(DFlashDraftModel):
352
+ """DSpark draft network: DFlash backbone + Markov / confidence heads."""
353
+
354
+ config_class = DSparkConfig
355
+
356
+ def __init__(self, config) -> None:
357
+ super().__init__(config)
358
+
359
+ self.markov_rank = int(getattr(config, "markov_rank", 0))
360
+ self.confidence_head_with_markov = bool(
361
+ getattr(config, "confidence_head_with_markov", True)
362
+ )
363
+
364
+ self.markov_head = build_markov_head(config)
365
+
366
+ self.confidence_head: Optional[nn.Module] = None
367
+ if getattr(config, "enable_confidence_head", False):
368
+ conf_input_dim = config.hidden_size
369
+ if self.confidence_head_with_markov:
370
+ if self.markov_head is None:
371
+ raise ValueError(
372
+ "confidence_head_with_markov=True requires a Markov head "
373
+ "(markov_rank > 0)."
374
+ )
375
+ conf_input_dim += self.markov_rank
376
+ self.confidence_head = AcceptRatePredictor(conf_input_dim)
377
+
378
+ # DeepSpec builds the heads before post_init so they get the HF normal
379
+ # initializer (std=initializer_range). The base DFlash __init__ already
380
+ # ran post_init before these heads existed, so re-apply _init_weights to
381
+ # the heads to reproduce DeepSpec's initialization exactly.
382
+ if self.markov_head is not None:
383
+ self.markov_head.apply(self._init_weights)
384
+ if self.confidence_head is not None:
385
+ self.confidence_head.apply(self._init_weights)
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b815c610734268d1b84f926d515cbdc9a9ba461eb617cbbb10bfc65498ab6ad7
3
+ size 1305601530