# How to Implement DFlash with Custom Target Models: A Complete Integration Guide

> Integrate DFlash with custom target models using this complete guide. Learn essential steps like matching hidden sizes and verifying target attributes for seamless implementation.

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

---

**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`](https://github.com/z-lab/dflash/blob/main/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 at `model.embed_tokens` or discoverable via the helper logic in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py) that traverses `model.model.embed_tokens` or `model.language_model.embed_tokens`.
- **`lm_head`**: The language modeling head. If absent, the code falls back to `embed_tokens.as_linear` or the embedding matrix transpose.
- **`layers`**: A list of transformer layers accessible via `model.layers`, `model.model.layers`, or `model.language_model.layers`.

## Core DFlash Components in model.py

The speculative decoding algorithm resides in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/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()`:

```python
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 HuggingFace `AutoModelForCausalLM`.
- If your target nests embeddings deeply, the helper in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py) traverses common paths to locate `embed_tokens`.
- Ensure the draft's `target_layer_ids` (computed in `dflash/model.py:27-36`) reference valid layer indices in your target.

### vLLM Backend

For vLLM servers, enable DFlash via the `--speculative-config` JSON:

```bash
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/MyCustomLLM` can be any HuggingFace model compatible with vLLM's `AutoModelForCausalLM` loader.
- The draft model path must point to a DFlash checkpoint containing a `dflash_config` dictionary in its [`config.json`](https://github.com/z-lab/dflash/blob/main/config.json).
- `num_speculative_tokens` must be ≤ the draft's `block_size` defined in its configuration.

### SGLang Backend

SGLang uses command-line flags to configure DFlash speculative decoding:

```bash
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 DFLASH` enables the block-diffusion draft.
- `--speculative-draft-model-path` specifies your DFlash checkpoint.
- `--speculative-num-draft-tokens` sets the block size (must be ≤ draft's `block_size`).
- `--trust-remote-code` is 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`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py):

```python
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 in `dflash/model_mlx.py:65-89`) reads the draft's [`config.json`](https://github.com/z-lab/dflash/blob/main/config.json) and constructs a `DFlashDraftModel`.
- The target model must expose a Qwen-style interface with `embed_tokens`, `lm_head`, and `model.layers`.
- `_GDNStateCapture` in `dflash/model_mlx.py:24-34` handles 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_size` equals target's hidden size.
- **Layer accessibility**: Target exposes transformer layers at standard paths (`model.layers`, `model.model.layers`, or `model.language_model.layers`).
- **Embedding interface**: `embed_tokens` attribute exists or is discoverable via the helper in [`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py).
- **LM head availability**: `lm_head` exists or can fall back to embedding matrix operations.
- **Draft configuration**: Checkpoint contains `dflash_config` with valid `target_layer_ids` (auto-computed if absent) and `block_size` matching your desired `num_speculative_tokens`.
- **Backend-specific flags**: `trust_remote_code=True` for Transformers and SGLang; correct `--speculative-config` JSON for vLLM; MLX compatibility for Apple Silicon.

## Key Source Files Reference

- **[`dflash/model.py`](https://github.com/z-lab/dflash/blob/main/dflash/model.py)** – Core speculative algorithm, draft model class, and `spec_generate()` public API (`dflash/model.py:350-367`).  
  <https://github.com/z-lab/dflash/blob/main/dflash/model.py>

- **[`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py)** – MLX-specific draft implementation, KV-cache handling (`_GDNStateCapture` at `dflash/model_mlx.py:24-34`), and `stream_generate()`.  
  <https://github.com/z-lab/dflash/blob/main/dflash/model_mlx.py>

- **[`dflash/__init__.py`](https://github.com/z-lab/dflash/blob/main/dflash/__init__.py)** – Lazy imports exposing `DFlashDraftModel`, `extract_context_feature`, and `sample`.  
  <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_head` interfaces.
- Use **`spec_generate()`** (`dflash/model.py:350-367`) for direct Transformers integration with `trust_remote_code=True`.
- Configure **vLLM** and **SGLang** via command-line speculative flags, ensuring `num_speculative_tokens` does not exceed the draft's `block_size`.
- For **MLX** on Apple Silicon, use `load_draft()` and `stream_generate()` from [`dflash/model_mlx.py`](https://github.com/z-lab/dflash/blob/main/dflash/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`](https://github.com/z-lab/dflash/blob/main/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`](https://github.com/z-lab/dflash/blob/main/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`](https://github.com/z-lab/dflash/blob/main/config.json) specifying `target_layer_ids`, `hidden_size`, and `block_size`.