How to Implement DFlash with Custom Target Models: A Complete Integration Guide
TLDR: To implement DFlash with custom target models, ensure your draft model's hidden size matches the target's, verify the target exposes embed_tokens and lm_head attributes, and use the appropriate backend API—spec_generate() for Transformers, speculative configuration for vLLM/SGLang, or stream_generate() for MLX.
DFlash is a lightweight block-diffusion draft model designed to accelerate large language models through speculative decoding. When you implement DFlash with custom target models from the z-lab/dflash repository, you can integrate any HuggingFace-compatible architecture by satisfying minimal interface requirements and selecting the appropriate backend-specific implementation.
Architectural Requirements for Custom Target Models
Before integrating a custom target, verify that your model architecture meets the following interface requirements enforced by the core logic in dflash/model.py.
Hidden Size Compatibility
The draft model's config.hidden_size must exactly match the target model's hidden size. This is rigidly enforced because extract_context_feature in dflash/model.py:39-45 concatenates hidden states from the target layers directly into the draft's context tensor. If sizes differ, you must train a new DFlash draft model specifically for your target's hidden size.
Model Interface Requirements
Your custom target must expose the following attributes:
embed_tokens: The embedding layer, typically located atmodel.embed_tokensor discoverable via the helper logic indflash/model.pythat traversesmodel.model.embed_tokensormodel.language_model.embed_tokens.lm_head: The language modeling head. If absent, the code falls back toembed_tokens.as_linearor the embedding matrix transpose.layers: A list of transformer layers accessible viamodel.layers,model.model.layers, ormodel.language_model.layers.
Core DFlash Components in model.py
The speculative decoding algorithm resides in dflash/model.py. Understanding these key functions helps when debugging custom integrations:
build_target_layer_ids(dflash/model.py:27-36): Computes which hidden layers the draft will sample from based on the draft's depth and the target's total layer count.extract_context_feature(dflash/model.py:39-45): Concatenates selected hidden states into a single context tensor fed into the draft.DFlashDraftModel.__call__(dflash/model.py:102-122): Implements the forward pass that mixes noise embeddings with target context via dual-attention layers.spec_generate(dflash/model.py:350-367): The public API that orchestrates the speculative loop, KV cache management, and token verification.
Implementation by Backend
Transformers (PyTorch) Backend
For direct HuggingFace Transformers integration, load both models and call spec_generate():
from transformers import AutoModel, AutoModelForCausalLM, AutoTokenizer
# 1. Load the DFlash draft model
draft = AutoModel.from_pretrained(
"z-lab/Qwen3-8B-DFlash-b16", # Replace with your draft checkpoint
trust_remote_code=True, # Required for DFlash class
dtype="auto",
device_map="cuda:0"
).eval()
# 2. Load your custom target model
target = AutoModelForCausalLM.from_pretrained(
"my-org/MyCustomLLM", # Your custom model repository
dtype="auto",
device_map="cuda:0"
).eval()
# 3. Initialize tokenizer
tokenizer = AutoTokenizer.from_pretrained("my-org/MyCustomLLM")
# 4. Prepare input
messages = [{"role": "user", "content": "Explain the concept of diffusion models."}]
input_ids = tokenizer.apply_chat_template(
messages,
return_tensors="pt",
add_generation_prompt=True,
enable_thinking=False
).to(draft.device)
# 5. Run speculative generation
output = draft.spec_generate(
input_ids=input_ids,
max_new_tokens=512,
temperature=0.0,
target=target,
stop_token_ids=[tokenizer.eos_token_id]
)
# 6. Decode output
print(tokenizer.decode(output[0], skip_special_tokens=False))
Critical details for custom targets:
- The target must implement the standard
forward(input_ids, position_ids, past_key_values, ...)signature used by HuggingFaceAutoModelForCausalLM. - If your target nests embeddings deeply, the helper in
dflash/model.pytraverses common paths to locateembed_tokens. - Ensure the draft's
target_layer_ids(computed indflash/model.py:27-36) reference valid layer indices in your target.
vLLM Backend
For vLLM servers, enable DFlash via the --speculative-config JSON:
vllm serve my-org/MyCustomLLM \
--speculative-config '{
"method": "dflash",
"model": "z-lab/MyCustomDraft",
"num_speculative_tokens": 15
}' \
--attention-backend flash_attn \
--max-num-batched-tokens 32768
Configuration requirements:
my-org/MyCustomLLMcan be any HuggingFace model compatible with vLLM'sAutoModelForCausalLMloader.- The draft model path must point to a DFlash checkpoint containing a
dflash_configdictionary in itsconfig.json. num_speculative_tokensmust be ≤ the draft'sblock_sizedefined in its configuration.
SGLang Backend
SGLang uses command-line flags to configure DFlash speculative decoding:
export SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN=1
python -m sglang.launch_server \
--model-path my-org/MyCustomLLM \
--speculative-algorithm DFLASH \
--speculative-draft-model-path z-lab/MyCustomDraft \
--speculative-num-draft-tokens 16 \
--attention-backend trtllm_mha \
--speculative-draft-attention-backend fa4 \
--mem-fraction-static 0.75 \
--trust-remote-code
Key flags:
--speculative-algorithm DFLASHenables the block-diffusion draft.--speculative-draft-model-pathspecifies your DFlash checkpoint.--speculative-num-draft-tokenssets the block size (must be ≤ draft'sblock_size).--trust-remote-codeis required to load the DFlash model class.
MLX Backend (Apple Silicon)
For Apple Silicon devices, use the MLX-specific APIs in dflash/model_mlx.py:
from dflash.model_mlx import load, load_draft, stream_generate
# 1. Load target model (MLX implementation)
model, tokenizer = load("my-org/MyCustomLLM")
# 2. Load DFlash draft model
draft = load_draft(
"z-lab/MyCustomDraft",
sliding_window_size=None # Optional: bound KV history
)
# 3. Prepare prompt
messages = [{"role": "user", "content": "What is the capital of France?"}]
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=True
)
# 4. Stream generation with block diffusion
for response in stream_generate(
model, draft, tokenizer, prompt,
block_size=16, # Must match draft's block_size
max_tokens=256,
temperature=0.6
):
print(response.text, end="", flush=True)
print("\nDone.")
MLX-specific details:
load_draft(defined indflash/model_mlx.py:65-89) reads the draft'sconfig.jsonand constructs aDFlashDraftModel.- The target model must expose a Qwen-style interface with
embed_tokens,lm_head, andmodel.layers. _GDNStateCaptureindflash/model_mlx.py:24-34handles KV cache rollback for Gated-Delta networks.
Quick Checklist for Custom Target Validation
Use this checklist before deploying DFlash with your custom model:
- Hidden size match: Draft's
config.hidden_sizeequals target's hidden size. - Layer accessibility: Target exposes transformer layers at standard paths (
model.layers,model.model.layers, ormodel.language_model.layers). - Embedding interface:
embed_tokensattribute exists or is discoverable via the helper indflash/model.py. - LM head availability:
lm_headexists or can fall back to embedding matrix operations. - Draft configuration: Checkpoint contains
dflash_configwith validtarget_layer_ids(auto-computed if absent) andblock_sizematching your desirednum_speculative_tokens. - Backend-specific flags:
trust_remote_code=Truefor Transformers and SGLang; correct--speculative-configJSON for vLLM; MLX compatibility for Apple Silicon.
Key Source Files Reference
-
dflash/model.py– Core speculative algorithm, draft model class, andspec_generate()public API (dflash/model.py:350-367).
https://github.com/z-lab/dflash/blob/main/dflash/model.py -
dflash/model_mlx.py– MLX-specific draft implementation, KV-cache handling (_GDNStateCaptureatdflash/model_mlx.py:24-34), andstream_generate().
https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py -
dflash/__init__.py– Lazy imports exposingDFlashDraftModel,extract_context_feature, andsample.
https://github.com/z-lab/dflash/blob/main/dflash/__init__.py
Summary
- Implement DFlash with custom target models by ensuring architectural compatibility: matching hidden sizes and standard
embed_tokens/lm_headinterfaces. - Use
spec_generate()(dflash/model.py:350-367) for direct Transformers integration withtrust_remote_code=True. - Configure vLLM and SGLang via command-line speculative flags, ensuring
num_speculative_tokensdoes not exceed the draft'sblock_size. - For MLX on Apple Silicon, use
load_draft()andstream_generate()fromdflash/model_mlx.py, ensuring your target follows Qwen-style layer conventions. - Verify integration using the Quick Checklist to confirm layer accessibility and hidden size alignment before deployment.
Frequently Asked Questions
Can I use DFlash with a custom target model that has a different hidden size than the draft?
No, the hidden sizes must match exactly. The draft's config.hidden_size must equal the target model's hidden size because extract_context_feature in dflash/model.py:39-45 concatenates target hidden states directly into the draft's context tensor. If your target has a different size, you must train a custom DFlash draft model with the matching hidden size.
What if my custom target model doesn't expose embed_tokens directly?
The DFlash implementation includes helper logic in dflash/model.py that traverses common attribute paths including model.embed_tokens, model.model.embed_tokens, and model.language_model.embed_tokens. As long as your model follows standard HuggingFace nesting conventions, the draft will locate the embeddings automatically. For highly custom architectures, add a property accessor exposing embed_tokens at the root level.
How do I verify that DFlash is actually accelerating my custom target model?
Monitor the acceptance rate and tokens-per-second metrics. In the Transformers backend, enable debug logging to see acceptance statistics during spec_generate(). For vLLM and SGLang, check the server logs for speculative decoding metrics. Successful acceleration typically shows draft acceptance rates above 60%, yielding 1.5x to 2.5x throughput improvements over standard autoregressive generation.
Can I train my own DFlash draft model for a custom target architecture?
Yes, though the training code is not included in the core integration files. The draft architecture is defined in dflash/model.py (DFlashDraftModel class, lines 102-122) and uses standard transformer components. To train for a custom target, collect hidden states from the target layers identified by build_target_layer_ids (dflash/model.py:27-36) and train the draft to predict the next token block. The resulting checkpoint must include a dflash_config dictionary in config.json specifying target_layer_ids, hidden_size, and block_size.
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 →