bridge.training.post_training.dspark.heads#

DSpark draft heads: the sequential (Markov / RNN) head and the confidence head.

DSpark (arXiv:2607.05147) draft training. The parallel backbone predicts every position of a block in one pass; these small heads add the pieces the paper needs on top of that:

  • the sequential head injects intra-block token dependency onto the parallel logits, as a factorized P(x | x0) = prod_k p_k(x_k | x0, x_1..x_{k-1}) bias. VanillaMarkov is a memoryless low-rank token-to-token additive logit bias; GatedMarkovHead gates the previous-token embedding by the backbone hidden state; RNNHead carries a recurrent state so position k sees the full prefix.

  • the confidence head predicts a per-position acceptance logit.

Only the training-time (teacher-forced, parallel) apply_block_logits path is provided here; the serial autoregressive sampling used at inference is a separate concern. These modules depend only on torch.

Module Contents#

Classes#

VanillaMarkov

Memoryless low-rank token-to-token additive logit bias.

GatedMarkovHead

VanillaMarkov whose previous-token embedding is gated by the backbone hidden.

RNNHead

GRU-like head carrying recurrent state across a block so position k sees x_1..x_{k-1}.

ConfidenceHead

Per-position acceptance-logit predictor (a single linear projection to a scalar).

Functions#

build_markov_head

Construct the sequential head, or None when markov_rank == 0 (disabled).

API#

class bridge.training.post_training.dspark.heads.VanillaMarkov(*, vocab_size: int, markov_rank: int)#

Bases: torch.nn.Module

Memoryless low-rank token-to-token additive logit bias.

The step bias is markov_w2(markov_w1[prev_token]): a rank-markov_rank factorization of a [vocab, vocab] transition-bias matrix.

Initialization

get_prev_embeddings(token_ids: torch.Tensor) torch.Tensor#

Look up markov_w1 for token_ids (any shape) -> [..., markov_rank].

project_bias(latent_states: torch.Tensor) torch.Tensor#

Project [..., markov_rank] latent states to a [..., vocab] logit bias.

compute_step_bias(
token_ids: torch.Tensor,
hidden_states: torch.Tensor | None,
) torch.Tensor#

Logit bias for the previous tokens token_ids -> [..., vocab] (hidden unused).

apply_block_logits(
base_logits: torch.Tensor,
*,
token_ids: torch.Tensor,
hidden_states: torch.Tensor | None,
) torch.Tensor#

Add the (teacher-forced) sequential bias to the whole block’s logits.

Parameters:
  • base_logits – Backbone logits [batch, num_blocks, block_size, vocab].

  • token_ids – Teacher-forced previous token per position [batch, num_blocks, block_size].

  • hidden_states – Backbone hidden [batch, num_blocks, block_size, hidden] (used by the gated / RNN heads, ignored here).

Returns:

Corrected logits [batch, num_blocks, block_size, vocab].

class bridge.training.post_training.dspark.heads.GatedMarkovHead(
*,
vocab_size: int,
markov_rank: int,
hidden_size: int,
)#

Bases: bridge.training.post_training.dspark.heads.VanillaMarkov

VanillaMarkov whose previous-token embedding is gated by the backbone hidden.

Initialization

compute_step_bias(
token_ids: torch.Tensor,
hidden_states: torch.Tensor | None,
) torch.Tensor#

Gate markov_w1[prev] by sigmoid(gate_proj([hidden; prev])) before markov_w2.

Parameters:
  • token_ids – Previous token ids [..., ].

  • hidden_states – Backbone hidden [..., hidden] (required).

Returns:

Logit bias [..., vocab].

class bridge.training.post_training.dspark.heads.RNNHead(*, vocab_size: int, markov_rank: int, hidden_size: int)#

Bases: bridge.training.post_training.dspark.heads.VanillaMarkov

GRU-like head carrying recurrent state across a block so position k sees x_1..x_{k-1}.

Initialization

_rnn_step(
state: torch.Tensor,
prev_embeddings: torch.Tensor,
hidden_states: torch.Tensor,
) tuple[torch.Tensor, torch.Tensor]#

One recurrent step.

Parameters:
  • state – Previous recurrent state [..., markov_rank].

  • prev_embeddings – markov_w1[prev] [..., markov_rank].

  • hidden_states – Backbone hidden at this position [..., hidden].

Returns:

(new_state [..., markov_rank], logit_bias [..., vocab]).

apply_block_logits(
base_logits: torch.Tensor,
*,
token_ids: torch.Tensor,
hidden_states: torch.Tensor | None,
) torch.Tensor#

Unroll the recurrence over the block (teacher-forced). Layout matches the base class.

bridge.training.post_training.dspark.heads.build_markov_head(
*,
markov_rank: int,
markov_head_type: str,
vocab_size: int,
hidden_size: int,
) torch.nn.Module | None#

Construct the sequential head, or None when markov_rank == 0 (disabled).

Parameters:
  • markov_rank – Low-rank size; 0 disables the head.

  • markov_head_type – One of "vanilla", "gated", "rnn".

  • vocab_size – Draft vocabulary size.

  • hidden_size – Backbone hidden size (used by gated / rnn).

Returns:

The head module, or None.

class bridge.training.post_training.dspark.heads.ConfidenceHead(input_dim: int)#

Bases: torch.nn.Module

Per-position acceptance-logit predictor (a single linear projection to a scalar).

Initialization

forward(features: torch.Tensor) torch.Tensor#

Project [..., input_dim] features to a [...] acceptance logit.