# How DFlash Extracts Context Features for Speculative Generation

> Discover how DFlash extracts context features for speculative generation by sampling hidden states and creating a compact representation for your draft model.

- Repository: [Z Lab/dflash](https://github.com/z-lab/dflash)
- Tags: how-to-guide
- Published: 2026-04-17

---

**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`](https://github.com/z-lab/dflash/blob/main/dflash/model.py).

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

```python
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`](https://github.com/z-lab/dflash/blob/main/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`](https://github.com/z-lab/dflash/blob/main/dflash/model.py). This function transforms the raw hidden states from the target model into a unified context tensor.

```python
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:

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

```

This extraction (lines 99-100 in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/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:

```python
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`](https://github.com/z-lab/dflash/blob/main/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`.