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:
- Prompt tokens (
input_ids) – the tokenized input sequence - 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 dimensionnum_hidden_layers– number of layers in the draft (typically 4–8)target_layer_ids– the list of layer indices you are training againstblock_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:
- Select target layers using
build_target_layer_idsindflash/model.pyto 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, specifyingtarget_layer_ids,block_size, andmask_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_featureduring 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.jsonto enable loading viaload_draftfor 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 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →