# How to Train Custom DFlash Draft Models for Specific Target LLMs

> Learn to train custom DFlash draft models by selecting target LLM layers, configuring architecture, and using hidden states and token predictions for effective supervision.

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

---

**To train custom DFlash draft models, you select specific target layers from a base LLM, configure a lightweight draft architecture via `DFlashConfig`, and supervise training using both the target's hidden states (MSE loss) and token predictions (cross-entropy loss).**

DFlash is an open-source block-diffusion acceleration framework that trains compact draft models to predict the **hidden representations** of larger target models at specific layers. By training a draft to mimic these intermediate activations, you can generate speculative token blocks that the target model verifies in parallel, achieving significant inference speedups. This guide walks through the complete training workflow using the `z-lab/dflash` codebase.

## Understanding DFlash Draft Model Architecture

DFlash draft models are lightweight **block-diffusion** networks that learn to predict the hidden representations of a target LLM at a subset of its layers. Unlike traditional speculative decoding drafts that predict single tokens, DFlash generates blocks of tokens by conditioning on the target model's intermediate activations.

The draft architecture accepts two inputs during training:
1. **Prompt tokens** (`input_ids`) – the tokenized input sequence
2. **Target hidden states** – concatenated hidden representations from selected layers of the target model

The draft outputs logits for the next token block, which are trained against both the true next tokens (cross-entropy) and the target hidden states (MSE).

## Step 1: Select Target Layers and Build Layer IDs

The draft model does not recreate the entire target architecture; it only predicts a concatenation of hidden states from a small set of **target layers**. You must specify which layers to emulate via `target_layer_ids` in the configuration.

### Using build_target_layer_ids

The utility `build_target_layer_ids` in **[`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py)** helps compute these IDs based on the number of target layers you want to use:

```python
from dflash.model import build_target_layer_ids

# For a 32-layer target model, select 4 evenly spaced layers

target_layer_ids = build_target_layer_ids(
    num_target_layers=4,
    num_draft_layers=4,
)

# Returns: (0, 8, 16, 24) or similar evenly spaced indices

```

This function is defined at lines 27–34 in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py)【`dflash/model.py#L27-L34`](https://github.com/z-lab/dflash/blob/main/dflash/model.py#L27-L34).

## Step 2: Configure the Draft Model

The draft’s hyperparameters are stored in a `DFlashConfig` data class (MLX version) or equivalent in the PyTorch version. The config fields are populated from a JSON file that the `load_draft` helper reads.

### DFlashConfig Parameters

Key configuration fields include:

- **`hidden_size`** – must match the target model's hidden dimension
- **`num_hidden_layers`** – number of layers in the draft (typically 4–8)
- **`target_layer_ids`** – the list of layer indices you are training against
- **`block_size`** – the number of tokens the draft predicts in one forward pass (must match the target’s context window)
- **`mask_token_id`** – the token used for "blank" positions in the draft's output

In **[`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py)**, the `DFlashConfig` dataclass is defined at lines 28–46【`dflash/model_mlx.py#L28-L46`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py#L28-L46).

Example configuration for a Qwen3-8B target:

```python
import json
from pathlib import Path

config_dict = {
    "hidden_size": 4096,
    "num_hidden_layers": 4,
    "num_attention_heads": 32,
    "num_key_value_heads": 8,
    "head_dim": 128,
    "intermediate_size": 20480,
    "vocab_size": 151936,
    "rms_norm_eps": 1e-06,
    "rope_theta": 1000000.0,
    "max_position_embeddings": 32768,
    "block_size": 16,
    "dflash_config": {
        "target_layer_ids": [0, 8, 16, 24],
        "mask_token_id": 151643,
    },
    "num_target_layers": 4,
    "rope_scaling": None,
}

Path("draft_config.json").write_text(json.dumps(config_dict, indent=2))

```

## Step 3: Training Loop with Hidden State Supervision

The training loop follows the pattern used in the inference code. The crucial operation is extracting the target hidden features and supervising the draft to reproduce them.

### Extracting Target Hidden States

Use `extract_context_feature` from **[`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py)** to concatenate hidden states from your selected target layers:

```python
from dflash.model import extract_context_feature

# Assuming target model outputs hidden states

with torch.no_grad():
    out = target_model(input_ids, output_hidden_states=True, return_dict=True)
    target_hidden = extract_context_feature(
        out.hidden_states,      # Tuple of hidden states from all layers

        target_layer_ids        # Your selected layer indices, e.g., (0, 8, 16, 24)

    )

# target_hidden shape: (batch, seq_len, hidden_dim * num_target_layers)

```

This function is defined at lines 99–102 in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py)【`dflash/model.py#L99-L102`](https://github.com/z-lab/dflash/blob/main/dflash/model.py#L99-L102).

### Computing Training Losses

You can train with either or both of:

- **Mean-squared error (MSE)** between the draft's hidden projection and the target hidden tensor
- **Cross-entropy** on the draft's token logits against the true next tokens

The draft forward call accepts the target hidden tensor to condition its predictions:

```python

# Draft forward (PyTorch variant)

draft_logits = draft(
    input_ids,
    target_hidden=target_hidden,
    cache=None  # or your KV cache implementation

)

# Compute cross-entropy loss

loss_ce = torch.nn.functional.cross_entropy(
    draft_logits[:, :-1].reshape(-1, draft_logits.size(-1)),
    input_ids[:, 1:].reshape(-1),
    reduction='mean'
)

# Compute MSE on hidden projections (if your draft has fc projection)

loss_mse = torch.nn.functional.mse_loss(
    draft.fc(target_hidden),
    target_hidden.detach()
)

# Combined loss

alpha = 0.5  # weighting factor

loss = loss_ce + alpha * loss_mse
loss.backward()
optimizer.step()

```

For the MLX version, the forward signature is similar:

```python

# Draft forward (MLX variant)

draft_logits = draft(input_ids, target_hidden, draft.make_cache())

```

See the `__call__` implementation in [`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py) at lines 152–158【`dflash/model_mlx.py#L152-L158`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py#L152-L158).

## Training Examples

### PyTorch Training Skeleton

Complete training loop for a Qwen3-8B target with a 4-layer draft:

```python
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from dflash.model import DFlashDraftModel, extract_context_feature, build_target_layer_ids

# 1️⃣ Load target model

target = AutoModelForCausalLM.from_pretrained(
    "Qwen/Qwen3-8B", 
    torch_dtype=torch.bfloat16
).cuda()
target.eval()

# 2️⃣ Choose target layers (e.g., 4 layers evenly spaced)

target_layer_ids = build_target_layer_ids(
    num_target_layers=4,
    num_draft_layers=4,
)
target.config.dflash_config = {
    "target_layer_ids": target_layer_ids, 
    "mask_token_id": 151643
}

# 3️⃣ Create draft model (mirroring target dimensions)

draft_cfg = DFlashDraftModel.Config.from_target(
    target, 
    target_layer_ids, 
    block_size=16
)
draft = DFlashDraftModel(draft_cfg).cuda()
optimizer = torch.optim.AdamW(draft.parameters(), lr=1e-4)

# 4️⃣ Training loop

tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-8B")
for batch in dataloader:  # Your data loading logic

    input_ids = tokenizer(
        batch["text"], 
        return_tensors="pt"
    )["input_ids"].cuda()
    
    with torch.no_grad():
        # Run target and capture hidden states

        out = target(
            input_ids, 
            output_hidden_states=True, 
            return_dict=True
        )
        target_hidden = extract_context_feature(
            out.hidden_states, 
            target_layer_ids
        )

    # Draft forward

    draft_logits = draft(
        input_ids, 
        target_hidden=target_hidden, 
        cache=None
    )

    # Compute losses

    loss_ce = torch.nn.functional.cross_entropy(
        draft_logits[:, :-1].reshape(-1, draft_logits.size(-1)),
        input_ids[:, 1:].reshape(-1),
    )
    loss_mse = torch.nn.functional.mse_loss(
        draft.fc(target_hidden), 
        target_hidden.detach()
    )
    loss = loss_ce + 0.5 * loss_mse

    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

```

### MLX Training Skeleton (Apple Silicon)

For training on Apple Silicon using the MLX framework:

```python
import json, pathlib
import mlx.core as mx
import mlx.nn as nn
from dflash.model_mlx import DFlashDraftModel, DFlashConfig
from mlx_lm import load as mlx_load

# 1️⃣ Load target (MLX) model

target, tokenizer = mlx_load("Qwen/Qwen3-4B")
target.eval()

# 2️⃣ Choose layers to supervise

target_layer_ids = (2, 6, 10, 14)

# 3️⃣ Build draft config

cfg = DFlashConfig(
    hidden_size=target.config.hidden_size,
    num_hidden_layers=4,
    num_attention_heads=target.config.num_attention_heads,
    num_key_value_heads=target.config.num_key_value_heads,
    head_dim=target.config.head_dim,
    intermediate_size=target.config.intermediate_size,
    vocab_size=target.config.vocab_size,
    rms_norm_eps=target.config.rms_norm_eps,
    rope_theta=target.config.rope_theta,
    max_position_embeddings=target.config.max_position_embeddings,
    block_size=16,
    target_layer_ids=target_layer_ids,
    num_target_layers=len(target_layer_ids),
    mask_token_id=tokenizer.mask_token_id,
)

# 4️⃣ Instantiate draft

draft = DFlashDraftModel(cfg)

# 5️⃣ Optimiser

optim = mx.optim.Adam(draft.parameters(), lr=1e-4)

# 6️⃣ Training loop

for batch in dataset:
    tokenized = tokenizer.apply_chat_template(
        [{"role": "user", "content": batch}], 
        add_generation_prompt=True
    )
    inputs = mx.array(tokenized).reshape(1, -1)

    # Target forward (captures hidden states)

    out = target(inputs, output_hidden_states=True)
    target_hidden = mx.concat(
        [out.hidden_states[i] for i in target_layer_ids], 
        axis=-1
    )

    # Draft forward

    draft_logits = draft(inputs, target_hidden, draft.make_cache())

    # CE loss on next-token prediction

    loss_ce = mx.nn.losses.cross_entropy(
        draft_logits[:, :-1], 
        inputs[:, 1:]
    ).mean()

    # MSE on hidden projection

    hidden_pred = draft.fc(target_hidden)
    loss_mse = mx.mean((hidden_pred - target_hidden) ** 2)

    loss = loss_ce + 0.5 * loss_mse
    loss.backward()
    optim.step()
    optim.zero_grad()

```

## Saving and Deploying Your Trained Draft

After training, serialize the model to **Safetensors** and write a matching [`config.json`](https://github.com/z-lab/dflash/blob/main/config.json). The `load_draft` routine expects exactly this layout:

```python
import json
import pathlib
import mlx.core as mx

# After training finishes

save_dir = pathlib.Path("my_dflash_draft")
save_dir.mkdir(parents=True, exist_ok=True)

# 1️⃣ Save weights (safetensors format)

weights = {k: v for k, v in draft.named_parameters()}
mx.save(weights, save_dir / "model.safetensors")

# 2️⃣ Write config.json (mirrors DFlashConfig)

config_dict = {
    "hidden_size": draft.config.hidden_size,
    "num_hidden_layers": draft.config.num_hidden_layers,
    "num_attention_heads": draft.config.num_attention_heads,
    "num_key_value_heads": draft.config.num_key_value_heads,
    "head_dim": draft.config.head_dim,
    "intermediate_size": draft.config.intermediate_size,
    "vocab_size": draft.config.vocab_size,
    "rms_norm_eps": draft.config.rms_norm_eps,
    "rope_theta": draft.config.rope_theta,
    "max_position_embeddings": draft.config.max_position_embeddings,
    "block_size": draft.config.block_size,
    "dflash_config": {
        "target_layer_ids": list(draft.config.target_layer_ids),
        "mask_token_id": draft.config.mask_token_id,
    },
    "num_target_layers": draft.config.num_target_layers,
    "rope_scaling": draft.config.rope_scaling,
}
(save_dir / "config.json").write_text(json.dumps(config_dict, indent=2))

```

The resulting directory can be uploaded to Hugging Face and loaded with `load_draft` exactly as shown in the repository’s inference examples. The loading routine in [`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py) (lines 70–90) expects the [`config.json`](https://github.com/z-lab/dflash/blob/main/config.json) and `*.safetensors` files to be present【`dflash/model_mlx.py#L70-L90`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py#L70-L90).

## Summary

To train custom DFlash draft models for specific target models:

- **Select target layers** using `build_target_layer_ids` in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py) to determine which intermediate activations the draft will predict【`dflash/model.py#L27-L34`](https://github.com/z-lab/dflash/blob/main/dflash/model.py#L27-L34).
- **Configure the draft** via `DFlashConfig` (MLX) or equivalent PyTorch classes, specifying `target_layer_ids`, `block_size`, and `mask_token_id`【`dflash/model_mlx.py#L28-L46`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py#L28-L46).
- **Extract target hidden states** using `extract_context_feature` during the forward pass to capture the specific layer activations【`dflash/model.py#L99-L102`](https://github.com/z-lab/dflash/blob/main/dflash/model.py#L99-L102).
- **Train with dual supervision** combining **MSE loss** on hidden projections and **cross-entropy** on token logits to align the draft with both the target's representation space and output distribution.
- **Save to Safetensors** with a matching [`config.json`](https://github.com/z-lab/dflash/blob/main/config.json) to enable loading via `load_draft` for inference.

## Frequently Asked Questions

### What is the difference between DFlash and standard speculative decoding?

Standard speculative decoding uses a smaller standalone language model to draft tokens, which the target model then verifies. **DFlash** instead trains a **block-diffusion** draft that generates token blocks conditioned on the target model's own intermediate hidden states. This allows the draft to produce higher-quality speculations that are more aligned with the target's internal representations, leading to better acceptance rates and larger speedups.

### How do I choose which target layers to supervise?

Select target layers that capture high-level semantic features without requiring the full depth of the target model. In practice, **evenly spaced layers** work well. Use `build_target_layer_ids` in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py) to generate indices automatically based on your desired number of target layers and draft layers【`dflash/model.py#L27-L34`](https://github.com/z-lab/dflash/blob/main/dflash/model.py#L27-L34). For a 32-layer target and 4-layer draft, this might select layers 0, 8, 16, and 24.

### Can I train DFlash drafts on Apple Silicon?

Yes. The `z-lab/dflash` repository provides an MLX implementation in **[`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py)** specifically for Apple Silicon training and inference. Use `DFlashConfig` and `DFlashDraftModel` from this module, and replace PyTorch operations with MLX equivalents (`mx.array`, `mx.nn.losses`, etc.). The training logic remains identical: extract target hidden states via `output_hidden_states`, concatenate selected layers, and compute MSE and cross-entropy losses.

### What file format should I use to save the trained draft?

Save your trained draft using the **Safetensors** format for the weights and a JSON file for the configuration. The `load_draft` function in [`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py) expects a directory containing `model.safetensors` and [`config.json`](https://github.com/z-lab/dflash/blob/main/config.json)【`dflash/model_mlx.py#L70-L90`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py#L70-L90). Ensure your [`config.json`](https://github.com/z-lab/dflash/blob/main/config.json) includes a `dflash_config` key containing `target_layer_ids` and `mask_token_id`, along with standard transformer parameters like `hidden_size` and `num_hidden_layers`.