How DFlash Handles Temperature and Sampling During Text Generation

DFlash controls text generation randomness through a unified sample function that switches between greedy decoding at near-zero temperatures and probabilistic sampling via softmax-scaled logits at higher temperatures.

DFlash is a speculative decoding framework developed by the z-lab/dflash repository that accelerates large language model inference. Understanding how DFlash handles temperature and sampling during text generation is essential for controlling output randomness across both draft and target model stages.

Temperature-Controlled Sampling Logic

The core sampling mechanism resides in dflash/model.py, where a lightweight helper function determines whether to generate tokens deterministically or stochastically based on the temperature parameter.

The Core Sampling Function

The sample function (lines 48-55) implements the temperature scaling logic:

def sample(logits, temperature):
    if temperature < 1e-5:
        # Greedy decoding: select highest probability token

        return torch.argmax(logits, dim=-1)
    else:
        # Stochastic sampling: scale logits by temperature

        probs = torch.softmax(logits / temperature, dim=-1)
        return torch.multinomial(probs, num_samples=1).squeeze(-1)

When temperature falls below 1e-5, DFlash executes greedy decoding using torch.argmax. For any positive temperature value, the function scales the logits by dividing by the temperature parameter, applies softmax to obtain probabilities, and samples via torch.multinomial.

Greedy vs. Stochastic Decoding

The temperature parameter acts as a global scaling factor across the entire generation pipeline:

  • Temperature ≈ 0: Deterministic output; the model always selects the highest probability token
  • Temperature > 0: Stochastic output; lower values (0.1-0.5) produce focused, conservative text while higher values (0.8-1.2) increase diversity and creativity

How Temperature Flows Through the Generation Pipeline

DFlash propagates the temperature argument through the entire call chain: spec_generate → dflash_generate → sample. This ensures consistent sampling behavior across different generation stages.

First-Token Generation

After the initial prefill pass, DFlash generates the first new token using the same sample function (lines 96-99 in dflash/model.py):


# First token sampling

first_token = sample(logits, temperature)
output_ids = torch.cat([input_ids, first_token.unsqueeze(0)], dim=1)

The temperature supplied to dflash_generate directly influences this initial sampling step, ensuring that the randomness level is established from the first generated token.

Draft Model Sampling

When operating with block_size > 1, DFlash samples draft logits produced by the draft model. The draft tokens are generated using the identical sample(draft_logits, temperature) call before being validated against the target model's posterior.

This means both the smaller draft model and the larger target model use the same temperature-controlled sampling logic, maintaining coherence between speculative and verification phases.

Posterior Sampling

For each decoding step, the target model's logits undergo sampling again (lines 134-136):


# Posterior sampling after verification

new_token = sample(target_logits, temperature)
output_ids = torch.cat([output_ids, new_token.unsqueeze(0)], dim=1)

The sampled token becomes the final output for that position, and the acceptance length determines how many draft tokens are retained. Throughout this process, the temperature parameter remains constant, ensuring uniform randomness across the speculative decoding loop.

Practical Usage Examples

The following example demonstrates how to invoke DFlash generation with temperature control:

from transformers import AutoTokenizer, AutoModelForCausalLM
from dflash.model import DFlashDraftModel

# Load models

tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-1.8B")
target = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-1.8B")
draft = DFlashDraftModel.from_pretrained("Qwen/Qwen3-1.8B")

# Prepare input

prompt = "Explain quantum entanglement in simple terms."
input_ids = tokenizer(prompt, return_tensors="pt").input_ids

# Generate with temperature 0.7 (stochastic)

output_ids = draft.spec_generate(
    target=target,
    input_ids=input_ids,
    max_new_tokens=100,
    stop_token_ids=[tokenizer.eos_token_id],
    temperature=0.7,
)

print(tokenizer.decode(output_ids[0], skip_special_tokens=True))

To adjust randomness levels:


# Greedy decoding for deterministic output

greedy_output = draft.spec_generate(..., temperature=0.0)

# High temperature for creative, diverse generation

creative_output = draft.spec_generate(..., temperature=1.2)

Summary

  • Unified sampling function: DFlash uses a single sample function in dflash/model.py (lines 48-55) that switches between greedy argmax and temperature-scaled multinomial sampling based on the temperature threshold of 1e-5.

  • Temperature propagation: The temperature parameter flows through spec_generate → dflash_generate → sample, ensuring consistent randomness across first-token generation, draft model sampling, and posterior verification.

  • Deterministic vs. stochastic: Temperatures below 1e-5 trigger deterministic greedy decoding, while positive temperatures enable probabilistic sampling where lower values increase focus and higher values increase diversity.

Frequently Asked Questions

What happens when I set temperature to 0 in DFlash?

When you set temperature to 0 (or any value below 1e-5), DFlash executes greedy decoding in the sample function. The code path calls torch.argmax(logits, dim=-1) to select the highest probability token at each step, resulting in deterministic, reproducible output regardless of how many times you run the generation.

How does DFlash temperature sampling differ from standard Hugging Face transformers?

While both frameworks use temperature scaling (logits / temperature), DFlash implements this within a speculative decoding context where the same temperature must synchronize both draft and target model sampling. In dflash/model.py, the sample function is reused across draft token generation, first-token selection, and posterior verification, ensuring consistent randomness that standard transformers handles independently in their generate method.

Can I use different temperatures for draft and target models in DFlash?

According to the source code in dflash/model.py, DFlash uses a single temperature parameter that propagates through the entire call chain from spec_generate through dflash_generate to the sample function. There is no mechanism to specify separate temperatures for draft versus target model sampling; both models share the same temperature value to maintain coherence in the speculative decoding acceptance criteria.

Where is the sampling logic implemented in the DFlash codebase?

The core sampling logic resides in dflash/model.py at lines 48-55, within the sample helper function. This function handles both greedy decoding (temperature < 1e-5) and temperature-scaled multinomial sampling. The same file contains the generation loop at dflash_generate (starting around line 70) and the high-level API spec_generate, both of which invoke the sample function for token selection.

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 →