How to Train Custom DFlash Draft Models for Specific Target LLMs

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 helps compute these IDs based on the number of target layers you want to use:

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【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, 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:

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 to concatenate hidden states from your selected target layers:

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【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:


# 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:


# Draft forward (MLX variant)

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

See the __call__ implementation in 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:

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:

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. The load_draft routine expects exactly this layout:

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 (lines 70–90) expects the 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:

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 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 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 expects a directory containing model.safetensors and 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 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.

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 →