How DFlash Extracts Context Features for Speculative Generation

DFlash extracts context features by sampling hidden states from specific layers of the target model, skipping the embedding layer, and concatenating them along the feature dimension to create a compact representation for the draft model.

DFlash is a speculative decoding framework that accelerates text generation by using a small draft model to predict tokens, then verifying them with a large target model. To keep the draft model informed about the generation progress without recomputing the full target model for every token, DFlash builds a context representation from the target model's intermediate hidden states. This article explains exactly how this context feature extraction works according to the z-lab/dflash source code.

Selecting Target Layers for Context

Before generation begins, DFlash determines which layers of the target model will supply the context features. This configuration happens during the initialization of DFlashDraftModel in dflash/model.py.

The model computes target_layer_ids by reading from the configuration or using a heuristic spread across the model depth:

self.target_layer_ids = self.config.dflash_config.get(
    "target_layer_ids",
    build_target_layer_ids(config.num_target_layers, config.num_hidden_layers)
)

The build_target_layer_ids helper function (lines 27-36 in dflash/model.py) evenly distributes the selected layer IDs across the depth of the target model. This ensures the context captures information from different representational depths without overwhelming the draft model with the full hidden state dimension.

The Core Extraction Mechanism

The heart of the feature extraction pipeline is the extract_context_feature function defined at lines 39-45 in dflash/model.py. This function transforms the raw hidden states from the target model into a unified context tensor.

def extract_context_feature(hidden_states: list[torch.Tensor],
                           layer_ids: Optional[list[int]]) -> torch.Tensor:
    offset = 1                     # skip the embedding layer

    selected_states = [hidden_states[layer_id + offset] for layer_id in layer_ids]
    return torch.cat(selected_states, dim=-1)

Handling the Embedding Offset

Notice the offset of 1 applied when indexing into hidden_states. The target model returns hidden states beginning with the embedding layer output at index 0. Since the embedding representation lacks the deep contextual processing of transformer layers, DFlash skips this initial entry by adding the offset to each layer ID.

Feature Concatenation Strategy

The function selects tensors corresponding to the configured layer_ids and concatenates them along the feature dimension (dim=-1). This produces a single high-dimensional tensor that aggregates multi-layer contextual information. The resulting shape is (batch, seq_len, sum_hidden_sizes_of_selected_layers), giving the draft model access to rich, multi-scale representations of the already-generated text.

Integration in the Generation Loop

Context extraction happens at specific points during the speculative decoding cycle to ensure the draft model always operates on fresh, verified information.

Initial Pre-fill Context

After the target model processes the initial prompt (the pre-fill phase), DFlash extracts the first context representation:

target_hidden = extract_context_feature(
    output.hidden_states, model.target_layer_ids)

This extraction (lines 99-100 in dflash/model.py) provides the draft model with the starting context for generating the first block of speculative tokens.

Updating After Decoding Blocks

Following each speculative block, DFlash verifies the draft tokens against the target model and determines which tokens to accept. It then extracts a new context, but only for the accepted portion of the sequence:

target_hidden = extract_context_feature(
    output.hidden_states, model.target_layer_ids)[:, :acceptance_length + 1, :]

This slicing operation (lines 42-44 in the generation loop) ensures the draft model receives updated context reflecting only the verified tokens, avoiding contamination from rejected speculative tokens while maintaining computational efficiency.

Summary

  • Layer Selection: DFlash uses build_target_layer_ids to choose specific transformer layers evenly distributed across the target model depth, configurable via target_layer_ids in the model config.
  • Offset Handling: The extract_context_feature function applies an offset of 1 to skip the embedding layer output, focusing only on processed hidden states.
  • Concatenation: Selected hidden states are concatenated along dim=-1 to create a unified context tensor with shape (batch, seq_len, combined_features).
  • Timing: Context extraction occurs after the initial pre-fill and after each accepted block during decoding, with slicing restricted to [:acceptance_length + 1] to include only verified tokens.
  • Source Location: The implementation resides in dflash/model.py, specifically in the extract_context_feature function and the DFlashDraftModel initialization logic.

Frequently Asked Questions

What is the purpose of context feature extraction in DFlash?

Context feature extraction provides the draft model with a compact, information-rich representation of the tokens generated so far. Since the draft model is much smaller than the target model, it cannot maintain the same level of contextual understanding on its own. By feeding it concatenated hidden states from deep layers of the target model, DFlash enables accurate speculative generation without requiring the draft model to recompute the full forward pass for every context update.

Why does DFlash skip the first hidden state with an offset of 1?

The first entry in the hidden_states list corresponds to the embedding layer output, which represents raw token embeddings without contextual processing from transformer attention mechanisms. By adding an offset = 1, DFlash ensures it extracts only from transformer layers that contain contextualized representations, providing the draft model with semantically meaningful features rather than static embeddings.

How does DFlash handle context updates when tokens are rejected?

During speculative decoding, the target model verifies draft tokens and determines an acceptance_length. DFlash extracts new context features using slicing [:, :acceptance_length + 1, :] to include only the accepted tokens plus one position. This ensures the draft model never receives hidden state information derived from rejected tokens, maintaining generation quality while updating the context efficiently.

Can I configure which target layers provide the context features?

Yes, layer selection is fully configurable through the dflash_config dictionary in the model configuration. You can explicitly set target_layer_ids to specify exact layer indices, or rely on the default build_target_layer_ids heuristic which evenly spaces layer selections across the target model's depth based on num_target_layers and num_hidden_layers.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →