# Target Layer ID Mapping in DFlash: Connecting Draft and Target Models

> Unlock efficient feature extraction with DFlash target layer ID mapping. Connect draft and target models for streamlined speculative generation. Learn how it works.

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

---

**Target layer ID mapping in DFlash creates a deterministic bridge that maps specific layers of the full target model to each layer of the lightweight draft model, allowing efficient feature extraction during speculative generation.**

DFlash implements speculative decoding by running a full-sized target model alongside a lightweight draft model. The `z-lab/dflash` repository defines a deterministic layer mapping system that enables the draft model to leverage hidden states from selected layers of the target model, avoiding the computational overhead of processing every layer while maintaining generation quality.

## How Target Layer ID Mapping Works in DFlash

DFlash operates two transformer models simultaneously during text generation. The **target model** serves as the full-size teacher producing accurate hidden states, while the **draft model** acts as a fast student generating candidate tokens. Rather than accessing every layer of the target model, the draft utilizes a **target layer ID mapping**—a curated subset of layer indices that define which hidden states to extract and concatenate for context features.

### The Layer Selection Strategy

The mapping strategically skips the initial embedding layer and final LM head layers, focusing instead on intermediate representations. This selection ensures the draft model receives semantically rich features without processing redundant shallow or deep transformer layers.

## Building the Mapping with build_target_layer_ids

The core logic resides in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py), specifically within the `build_target_layer_ids` function. This utility calculates evenly spaced layer indices across the target model's depth.

```python
def build_target_layer_ids(num_target_layers: int, num_draft_layers: int):
    if num_draft_layers == 1:
        return [num_target_layers // 2]
    start = 1
    end = num_target_layers - 3
    span = end - start
    return [
        int(round(start + (i * span) / (num_draft_layers - 1)))
        for i in range(num_draft_layers)
    ]

```

The algorithm spreads draft-layer indices uniformly across the target depth, reserving the first layer (index 0) for embeddings and the final layers for the language modeling head. For a 32-layer target model and 4-layer draft model, this generates indices like `[2, 12, 22, 30]`.

## Storing and Accessing the Mapping

When instantiating a `DFlashDraftModel`, the mapping loads from configuration or generates dynamically:

```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)
)

```

This list persists throughout the generation pipeline, enabling consistent feature extraction across forward passes.

## Extracting Context Features During Generation

### The extract_context_feature Function

During generation, DFlash calls `extract_context_feature` (defined in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py) lines 39-45) to gather hidden states from the target model:

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

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

```

The function applies a fixed `offset` of 1 to align with the Transformer's output format, where the first entry represents embeddings. It concatenates the selected layer tensors along the hidden dimension, creating a composite context feature.

### Integration in the Generation Loop

The mapping activates during speculative generation blocks. As shown in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py) (lines 98-100 and 143-144):

```python
if block_size > 1:
    target_hidden = extract_context_feature(output.hidden_states, model.target_layer_ids)

```

After token acceptance, the context updates using the same mapping truncated to the acceptance length:

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

```

## MLX Backend Implementation

The target layer ID mapping extends to the MLX backend in [`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py). When loading a draft model, the system reads the saved configuration (lines 84-86):

```python
target_layer_ids=tuple(cfg["dflash_config"]["target_layer_ids"]),

```

The `_patch_model` function (lines 18-25) subsequently modifies the model to capture hidden states specifically at these indices, ensuring cross-backend consistency.

## Summary

- **Target layer ID mapping** determines which specific layers of the full target model supply hidden states to the draft model during DFlash speculative decoding.
- **`build_target_layer_ids`** in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py) generates uniformly spaced layer indices, excluding embeddings and final LM head layers.
- **`extract_context_feature`** concatenates hidden states from the mapped indices, applying an offset of 1 to skip embedding outputs.
- The mapping persists in `DFlashDraftModel.target_layer_ids` and functions identically across PyTorch and MLX backends.

## Frequently Asked Questions

### How does DFlash decide which target layers to map?

DFlash uses a deterministic spacing algorithm in `build_target_layer_ids` that calculates evenly distributed indices across the target model's depth. It reserves the first layer for embeddings and the final three layers for the LM head, selecting intermediate layers that provide optimal semantic features for the draft model.

### Can I customize the target layer IDs manually?

Yes. You can specify custom `target_layer_ids` in the `dflash_config` section of your model configuration. If not provided, the system automatically generates the mapping using `build_target_layer_ids` based on the target and draft layer counts.

### Why does extract_context_feature use an offset of 1?

The offset accounts for the Transformer's hidden states list structure, where index 0 contains the input embeddings. Adding 1 aligns the target layer IDs with the actual transformer layer outputs, ensuring the draft model receives meaningful hidden representations rather than raw embeddings.

### Does the mapping work the same way in the MLX backend?

Yes. The MLX implementation in [`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py) loads the same `target_layer_ids` from the configuration JSON and uses `_patch_model` to intercept hidden states at the specified indices, maintaining consistency with the PyTorch implementation.