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

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, specifically within the build_target_layer_ids function. This utility calculates evenly spaced layer indices across the target model's depth.

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:

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 lines 39-45) to gather hidden states from the target model:

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 (lines 98-100 and 143-144):

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:

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. When loading a draft model, the system reads the saved configuration (lines 84-86):

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 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 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.

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 →