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.VanillaMarkovis a memoryless low-rank token-to-token additive logit bias;GatedMarkovHeadgates the previous-token embedding by the backbone hidden state;RNNHeadcarries 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#
Memoryless low-rank token-to-token additive logit bias. |
|
|
|
GRU-like head carrying recurrent state across a block so position k sees |
|
Per-position acceptance-logit predictor (a single linear projection to a scalar). |
Functions#
Construct the sequential head, or |
API#
- class bridge.training.post_training.dspark.heads.VanillaMarkov(*, vocab_size: int, markov_rank: int)#
Bases:
torch.nn.ModuleMemoryless low-rank token-to-token additive logit bias.
The step bias is
markov_w2(markov_w1[prev_token]): a rank-markov_rankfactorization of a[vocab, vocab]transition-bias matrix.Initialization
- get_prev_embeddings(token_ids: torch.Tensor) torch.Tensor#
Look up
markov_w1fortoken_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,
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,
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.VanillaMarkovVanillaMarkovwhose previous-token embedding is gated by the backbone hidden.Initialization
- compute_step_bias(
- token_ids: torch.Tensor,
- hidden_states: torch.Tensor | None,
Gate
markov_w1[prev]bysigmoid(gate_proj([hidden; prev]))beforemarkov_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.VanillaMarkovGRU-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,
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,
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,
Construct the sequential head, or
Nonewhenmarkov_rank == 0(disabled).- Parameters:
markov_rank – Low-rank size;
0disables 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.ModulePer-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.