# coding=utf-8 """DSpark draft model: DFlash backbone + EAGLE-style Markov and confidence heads. DSpark shares SpecForge's DFlash block-diffusion drafter (dual-source KV injection via :class:`DFlashDraftModel`, anchor sampling, MASK-token noise stream) and adds two heads on top: - Markov head: a learned low-rank bias added to the draft logits, conditioned on the (teacher-forced) previous token. Three variants are supported, exactly mirroring DeepSpec: ``vanilla`` (memoryless bigram), ``gated`` (token-gated), and ``rnn`` (recurrent state across within-block positions). - Confidence head (AcceptRatePredictor): predicts a per-draft-position acceptance probability, trained against the empirical draft-vs-target accept rate (used at inference for adaptive block length). The Markov / confidence / accept-rate modeling is ported to match DeepSeek's DeepSpec one-for-one (``deepspec/modeling/dspark/{markov_head,common}.py``, MIT License). SpecForge structural differences (load-bearing): - There is no ``DFlashConfig``; SpecForge's :class:`DFlashDraftModel` uses a plain ``Qwen3Config`` plus a ``config.dflash_config`` dict. So :class:`DSparkConfig` subclasses ``Qwen3Config`` and declares the DSpark fields as top-level attributes; DFlash-carried fields (``block_size``, ``num_target_layers``, ``dflash_config``) stay as before. - The draft model has no ``embed_tokens`` / ``lm_head`` of its own (they live on the target and are passed into the online wrapper). The heads only depend on ``config.hidden_size`` / ``config.vocab_size``, so this does not matter for construction. - DeepSpec builds the heads *before* ``post_init`` so the HF initializer (normal, std=initializer_range) covers them. SpecForge's base ``__init__`` runs ``post_init`` before the DSpark heads exist, so we re-apply ``_init_weights`` to the heads here to reproduce DeepSpec's initialization exactly (without this, ``markov_w1`` would keep the nn.Embedding default N(0,1) and the Markov bias would be huge at init). """ from typing import Optional import torch import torch.nn as nn from transformers.models.qwen3.modeling_qwen3 import Qwen3Config from specforge.modeling.draft.dflash import DFlashDraftModel def _sample_tokens(logits: torch.Tensor, temperature: float = 0.0) -> torch.Tensor: """Greedy (temperature<1e-5) or multinomial sampling over the last dim. Mirrors DeepSpec ``deepspec/utils/sampling.py::sample_tokens``. Used only by the heads' inference-time ``sample_block_tokens`` (not the training forward). """ if temperature < 1e-5: return torch.argmax(logits, dim=-1) bsz, seq_len, vocab_size = logits.shape flat_logits = logits.reshape(-1, vocab_size) / temperature probs = torch.softmax(flat_logits, dim=-1) return torch.multinomial(probs, num_samples=1).reshape(bsz, seq_len) class DSparkConfig(Qwen3Config): """Configuration for the DSpark draft model. Extends ``Qwen3Config``. DSpark-specific fields are declared here; the DFlash-carried fields (``block_size``, ``num_target_layers``, and the nested ``dflash_config`` dict holding ``target_layer_ids`` / ``mask_token_id``) are consumed by the :class:`DFlashDraftModel` base ``__init__`` and must be present on the config object before constructing the model. """ model_type = "dspark" def __init__( self, markov_rank: int = 256, markov_head_type: str = "vanilla", enable_confidence_head: bool = True, confidence_head_with_markov: bool = True, **kwargs, ): super().__init__(**kwargs) self.markov_rank = markov_rank self.markov_head_type = markov_head_type self.enable_confidence_head = enable_confidence_head self.confidence_head_with_markov = confidence_head_with_markov class VanillaMarkov(nn.Module): """Memoryless low-rank learned bigram bias added to the draft logits. Ported from DeepSpec ``deepspec/modeling/dspark/markov_head.py``. """ 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: del hidden_states return self.project_bias(self.get_prev_embeddings(token_ids)) def apply_step_logits( self, logits: torch.Tensor, *, token_ids: torch.Tensor, hidden_states: Optional[torch.Tensor] = None, ) -> torch.Tensor: return logits + self.compute_step_bias(token_ids, hidden_states) def apply_block_logits( self, base_logits: torch.Tensor, *, token_ids: torch.Tensor, hidden_states: Optional[torch.Tensor] = None, ) -> torch.Tensor: if base_logits.size(2) == 0: return base_logits return base_logits + self.compute_step_bias(token_ids, hidden_states) def sample_block_tokens( self, base_logits: torch.Tensor, *, first_prev_token_ids: torch.Tensor, hidden_states: Optional[torch.Tensor] = None, temperature: float = 0.0, ): batch_size, proposal_len = base_logits.shape[:2] if proposal_len == 0: empty_tokens = torch.empty( batch_size, 0, dtype=torch.long, device=base_logits.device ) return empty_tokens, base_logits sampled_tokens = [] corrected_logits = [] prev_token_ids = first_prev_token_ids.long() for step_idx in range(proposal_len): step_hidden = ( None if hidden_states is None else hidden_states[:, step_idx, ...] ) step_logits = self.apply_step_logits( base_logits[:, step_idx, :], token_ids=prev_token_ids, hidden_states=step_hidden, ) corrected_logits.append(step_logits.unsqueeze(1)) next_token_ids = _sample_tokens( step_logits.unsqueeze(1), temperature=temperature ).squeeze(1) 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): """Token-gated Markov head (DeepSpec ``gated``). The previous-token embedding is gated by a sigmoid of [hidden; prev_emb] before projection, letting the backbone hidden modulate the bigram bias. """ 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: Optional[torch.Tensor], ) -> torch.Tensor: assert hidden_states is not None prev_embeddings = self.get_prev_embeddings(token_ids) gate_inputs = torch.cat([hidden_states, prev_embeddings], 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: 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) class RNNHead(VanillaMarkov): """Recurrent Markov head (DeepSpec ``rnn``). Maintains a GRU-like recurrent state across within-block positions, so position k can access the full prefix history x_{ [gate; candidate; output] self.joint_proj = nn.Linear(2 * markov_rank + hidden_size, 3 * markov_rank) def _rnn_step( self, state: torch.Tensor, prev_embeddings: torch.Tensor, hidden_states: torch.Tensor, ): z = torch.cat([state, prev_embeddings, hidden_states], dim=-1) proj = self.joint_proj(z) gate_raw, candidate_raw, output_raw = proj.chunk(3, dim=-1) gate = torch.sigmoid(gate_raw) candidate = torch.tanh(candidate_raw) new_state = gate * state + (1.0 - gate) * candidate bias = self.project_bias(torch.tanh(output_raw)) return new_state, bias def compute_step_bias( self, token_ids: torch.Tensor, hidden_states: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Stateless single-step bias (state initialized to zero).""" assert hidden_states is not None prev_embeddings = self.get_prev_embeddings(token_ids) state = torch.zeros_like(prev_embeddings) _, bias = self._rnn_step(state, prev_embeddings, hidden_states) return bias def apply_block_logits( self, base_logits: torch.Tensor, *, token_ids: torch.Tensor, hidden_states: Optional[torch.Tensor] = None, ) -> torch.Tensor: assert hidden_states is not None block_size = base_logits.size(-2) if block_size == 0: return base_logits leading_shape = base_logits.shape[:-2] state = torch.zeros( *leading_shape, self.markov_rank, device=base_logits.device, dtype=hidden_states.dtype, ) output_logits = [] for k in range(block_size): prev_emb = self.get_prev_embeddings(token_ids[..., k]) h_k = hidden_states[..., k, :] state, bias = self._rnn_step(state, prev_emb, h_k) output_logits.append(base_logits[..., k, :] + bias) return torch.stack(output_logits, dim=-2) def sample_block_tokens( self, base_logits: torch.Tensor, *, first_prev_token_ids: torch.Tensor, hidden_states: Optional[torch.Tensor] = None, temperature: float = 0.0, ): assert hidden_states is not None batch_size, proposal_len = base_logits.shape[:2] if proposal_len == 0: empty_tokens = torch.empty( batch_size, 0, dtype=torch.long, device=base_logits.device ) return empty_tokens, base_logits state = torch.zeros( batch_size, self.markov_rank, device=base_logits.device, dtype=hidden_states.dtype, ) sampled_tokens = [] corrected_logits = [] prev_token_ids = first_prev_token_ids.long() for step_idx in range(proposal_len): prev_emb = self.get_prev_embeddings(prev_token_ids) h_k = hidden_states[:, step_idx, :] state, bias = self._rnn_step(state, prev_emb, h_k) step_logits = base_logits[:, step_idx, :] + bias corrected_logits.append(step_logits.unsqueeze(1)) next_token_ids = _sample_tokens( step_logits.unsqueeze(1), temperature=temperature ).squeeze(1) 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 AcceptRatePredictor(nn.Module): """Per-position acceptance-probability predictor (a single linear head). Ported from DeepSpec ``deepspec/modeling/dspark/common.py``. """ def __init__(self, input_dim: int): super().__init__() self.proj = nn.Linear(int(input_dim), 1) def forward(self, features: torch.Tensor) -> torch.Tensor: return self.proj(features).squeeze(-1) def build_markov_head(config) -> Optional[nn.Module]: markov_rank = int(getattr(config, "markov_rank", 0)) assert markov_rank >= 0, f"markov_rank must be >= 0, got {markov_rank}" if markov_rank == 0: return None markov_head_type = str(getattr(config, "markov_head_type", "vanilla")).lower() if markov_head_type == "vanilla": return VanillaMarkov(vocab_size=config.vocab_size, markov_rank=markov_rank) if markov_head_type == "gated": return GatedMarkovHead( vocab_size=config.vocab_size, markov_rank=markov_rank, hidden_size=config.hidden_size, ) if markov_head_type == "rnn": return RNNHead( vocab_size=config.vocab_size, markov_rank=markov_rank, hidden_size=config.hidden_size, ) raise AssertionError(f"Unsupported markov_head_type: {markov_head_type!r}") class DSparkDraftModel(DFlashDraftModel): """DSpark draft network: DFlash backbone + Markov / confidence heads.""" config_class = DSparkConfig def __init__(self, config) -> None: super().__init__(config) self.markov_rank = int(getattr(config, "markov_rank", 0)) self.confidence_head_with_markov = bool( getattr(config, "confidence_head_with_markov", True) ) self.markov_head = build_markov_head(config) self.confidence_head: Optional[nn.Module] = None if getattr(config, "enable_confidence_head", False): conf_input_dim = config.hidden_size if self.confidence_head_with_markov: if self.markov_head is None: raise ValueError( "confidence_head_with_markov=True requires a Markov head " "(markov_rank > 0)." ) conf_input_dim += self.markov_rank self.confidence_head = AcceptRatePredictor(conf_input_dim) # DeepSpec builds the heads before post_init so they get the HF normal # initializer (std=initializer_range). The base DFlash __init__ already # ran post_init before these heads existed, so re-apply _init_weights to # the heads to reproduce DeepSpec's initialization exactly. if self.markov_head is not None: self.markov_head.apply(self._init_weights) if self.confidence_head is not None: self.confidence_head.apply(self._init_weights)