From e3ceccf78150a80ee058c4365cc36be9d4001cb2 Mon Sep 17 00:00:00 2001
From: Chenchen Hong <76396533+Hayden727@users.noreply.github.com>
Date: Sat, 11 Jul 2026 11:05:05 +0800
Subject: [PATCH] [dLLM] Make FDFO a framework capability for all dLLM
algorithms (#27551)
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
---
.../advanced_features/server_arguments.mdx | 6 +
.../diffusion_language_models.mdx | 13 ++
python/sglang/srt/dllm/algorithm/base.py | 121 +++++++++++-
.../srt/dllm/algorithm/joint_threshold.py | 174 ++++++++----------
.../srt/dllm/algorithm/low_confidence.py | 113 ++++--------
python/sglang/srt/dllm/config.py | 3 +
python/sglang/srt/dllm/mixin/req.py | 17 ++
python/sglang/srt/dllm/mixin/scheduler.py | 83 +++++++--
python/sglang/srt/managers/scheduler.py | 7 +-
python/sglang/srt/managers/tp_worker.py | 23 ++-
python/sglang/srt/managers/utils.py | 7 +-
python/sglang/srt/server_args.py | 7 +
...da2_mini.py => test_dllm_batching_fdfo.py} | 9 +-
.../dllm/test_dllm_fdfo_joint_threshold.py | 89 +++++++++
14 files changed, 468 insertions(+), 204 deletions(-)
rename test/registered/dllm/{test_llada2_mini.py => test_dllm_batching_fdfo.py} (89%)
create mode 100644 test/registered/dllm/test_dllm_fdfo_joint_threshold.py
diff --git a/docs_new/docs/advanced_features/server_arguments.mdx b/docs_new/docs/advanced_features/server_arguments.mdx
index 298dd0cf0..dbc9a4d09 100644
--- a/docs_new/docs/advanced_features/server_arguments.mdx
+++ b/docs_new/docs/advanced_features/server_arguments.mdx
@@ -2217,6 +2217,12 @@ Please consult the documentation below and [server_args.py](https://github.com/s
`None` |
Type: str |
+
+ | `--dllm-fdfo` |
+ First-Done-First-Out (FDFO) scheduling lets completed requests leave the batch immediately instead of waiting for slower requests, eliminating head-of-line blocking. Enabled by default; pass `--no-dllm-fdfo` to fall back to synchronous lockstep scheduling. Works with any dLLM algorithm. |
+ `True` |
+ Type: bool |
+
diff --git a/docs_new/docs/supported-models/diffusion_language_models.mdx b/docs_new/docs/supported-models/diffusion_language_models.mdx
index 1f31d952e..e704f7ae9 100644
--- a/docs_new/docs/supported-models/diffusion_language_models.mdx
+++ b/docs_new/docs/supported-models/diffusion_language_models.mdx
@@ -16,6 +16,19 @@ python3 -m sglang.launch_server \
--port 30000
```
+## First-Done-First-Out (FDFO) Scheduling
+
+FDFO scheduling is **enabled by default**: each request leaves the batch as soon as its block is resolved, instead of advancing in lockstep where fast-converging requests must wait for slow long-tail requests before leaving the batch (head-of-line blocking). This improves throughput and is orthogonal to `--dllm-algorithm`, so it works with any dLLM algorithm. Pass `--no-dllm-fdfo` to fall back to synchronous lockstep scheduling:
+
+```bash Command
+python3 -m sglang.launch_server \
+ --model-path inclusionAI/LLaDA2.0-mini \
+ --dllm-algorithm LowConfidence \
+ --no-dllm-fdfo \
+ --host 0.0.0.0 \
+ --port 30000
+```
+
## Example Configuration File
Depending on the algorithm selected, the configuration parameters vary.
diff --git a/python/sglang/srt/dllm/algorithm/base.py b/python/sglang/srt/dllm/algorithm/base.py
index 349ddf4cd..5272c518b 100644
--- a/python/sglang/srt/dllm/algorithm/base.py
+++ b/python/sglang/srt/dllm/algorithm/base.py
@@ -1,18 +1,131 @@
+from __future__ import annotations
+
+from typing import Any, List, Optional, Tuple, Union
+
+import torch
+
from sglang.srt.dllm.algorithm import get_algorithm
from sglang.srt.dllm.config import DllmConfig
+from sglang.srt.layers.logits_processor import LogitsProcessorOutput
+from sglang.srt.model_executor.forward_batch_info import ForwardBatch
+from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.server_args import ServerArgs
+DllmRunOutput = Tuple[
+ Union[LogitsProcessorOutput, torch.Tensor],
+ List,
+ Optional[List[int]],
+ Optional[List[Any]],
+ bool,
+]
+
class DllmAlgorithm:
+ """dLLM algorithm: subclasses implement ``step``; the base owns the
+ synchronous and FDFO (``--dllm-fdfo``) execution loops in ``run``.
+ """
- def __init__(
- self,
- config: DllmConfig,
- ):
+ def __init__(self, config: DllmConfig):
self.block_size = config.block_size
self.mask_id = config.mask_id
+ self.fdfo = config.first_done_first_out_mode
@staticmethod
def from_server_args(server_args: ServerArgs):
config = DllmConfig.from_server_args(server_args)
return get_algorithm(config)
+
+ def init_step_state(self, forward_batch: ForwardBatch) -> List[Any]:
+ return [None] * forward_batch.batch_size
+
+ def max_steps(self, block_size: int) -> int:
+ return block_size + 1
+
+ def step(
+ self,
+ forward_batch: ForwardBatch,
+ full_logits: torch.Tensor,
+ states: List[Any],
+ ) -> List[bool]:
+ """One denoise step, advancing ``forward_batch.input_ids``/``states`` in
+ place. Returns, per block, whether it was already complete *on entry* --
+ i.e. this forward persisted its final KV cache and it can be emitted.
+ """
+ raise NotImplementedError
+
+ def run(
+ self,
+ model_runner: ModelRunner,
+ forward_batch: ForwardBatch,
+ algo_states: Optional[List[Any]] = None,
+ ) -> DllmRunOutput:
+ if self.fdfo:
+ return self._run_fdfo(model_runner, forward_batch, algo_states)
+ return self._run_sync(model_runner, forward_batch)
+
+ def _block_start_list(self, forward_batch: ForwardBatch) -> List[int]:
+ batch_size = forward_batch.batch_size
+ input_ids = forward_batch.input_ids.view(batch_size, self.block_size)
+ return (input_ids != self.mask_id).sum(dim=1).tolist()
+
+ def _run_sync(
+ self, model_runner: ModelRunner, forward_batch: ForwardBatch
+ ) -> DllmRunOutput:
+ batch_size = forward_batch.batch_size
+ start_list = self._block_start_list(forward_batch)
+
+ out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
+ # No mask to denoise: return empty so process_batch_result_dllm skips the
+ # stream branch (matches the pre-refactor behavior).
+ if all(start == self.block_size for start in start_list):
+ return out.logits_output, [], None, None, out.can_run_graph
+
+ states = self.init_step_state(forward_batch)
+ for _ in range(self.max_steps(self.block_size)):
+ done = self.step(forward_batch, out.logits_output.full_logits, states)
+ if all(done):
+ break
+ out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
+
+ next_token_ids = forward_batch.input_ids.view(batch_size, self.block_size)
+ next_token_ids_list = [
+ next_token_ids[i, start_list[i] :] for i in range(batch_size)
+ ]
+ return out.logits_output, next_token_ids_list, None, None, out.can_run_graph
+
+ def _run_fdfo(
+ self,
+ model_runner: ModelRunner,
+ forward_batch: ForwardBatch,
+ algo_states: Optional[List[Any]],
+ ) -> DllmRunOutput:
+ batch_size = forward_batch.batch_size
+
+ if algo_states is None:
+ algo_states = [None] * batch_size
+ fresh: Optional[List[Any]] = None
+ states: List[Any] = []
+ for i, carried in enumerate(algo_states):
+ if carried is None:
+ if fresh is None:
+ fresh = self.init_step_state(forward_batch)
+ states.append(fresh[i])
+ else:
+ states.append(carried)
+
+ out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
+ done = self.step(forward_batch, out.logits_output.full_logits, states)
+
+ accept_length_per_req_cpu = [self.block_size if d else 0 for d in done]
+ next_token_ids_list = forward_batch.input_ids.view(
+ batch_size, self.block_size
+ ).tolist()
+ states_out = [None if done[i] else states[i] for i in range(batch_size)]
+
+ return (
+ out.logits_output,
+ next_token_ids_list,
+ accept_length_per_req_cpu,
+ states_out,
+ out.can_run_graph,
+ )
diff --git a/python/sglang/srt/dllm/algorithm/joint_threshold.py b/python/sglang/srt/dllm/algorithm/joint_threshold.py
index a572fda38..6866996c4 100644
--- a/python/sglang/srt/dllm/algorithm/joint_threshold.py
+++ b/python/sglang/srt/dllm/algorithm/joint_threshold.py
@@ -1,20 +1,21 @@
+from typing import Any, List
+
import numpy as np
import torch
import torch.nn.functional as F
from sglang.srt.dllm.algorithm.base import DllmAlgorithm
from sglang.srt.dllm.config import DllmConfig
-from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
-from sglang.srt.model_executor.model_runner import ModelRunner
class JointThreshold(DllmAlgorithm):
+ """Joint-threshold denoising: mask-to-token (M2T) unmasking plus token-to-token
+ (T2T) edits, finishing on no-change or an exhausted edit budget. Stateful (edit
+ budget + prompt mask), carried across FDFO rounds via ``dllm_algo_state``.
+ """
- def __init__(
- self,
- config: DllmConfig,
- ):
+ def __init__(self, config: DllmConfig):
super().__init__(config)
self.threshold = config.algorithm_config.get("threshold", 0.5)
self.edit_threshold = config.algorithm_config.get("edit_threshold", 0)
@@ -23,117 +24,96 @@ class JointThreshold(DllmAlgorithm):
)
self.penalty_lambda = config.algorithm_config.get("penalty_lambda", 0)
- def run(
- self,
- model_runner: ModelRunner,
- forward_batch: ForwardBatch,
- ) -> tuple[LogitsProcessorOutput | torch.Tensor, torch.Tensor | None, bool]:
+ def max_steps(self, block_size: int) -> int:
+ return block_size + self.max_post_edit_steps + 1
+
+ def init_step_state(self, forward_batch: ForwardBatch) -> List[Any]:
batch_size = forward_batch.batch_size
- device = forward_batch.input_ids.device
+ input_ids = forward_batch.input_ids.view(batch_size, self.block_size)
+ # Built once as a GPU tensor and reused across steps (no per-step
+ # host/device transfer); the FDFO carry keeps it in-process.
+ prompt_mask = input_ids != self.mask_id
+ return [
+ {
+ "post_edit_steps": 0,
+ "finished": False,
+ "prompt_mask": prompt_mask[i],
+ }
+ for i in range(batch_size)
+ ]
- mask_index = forward_batch.input_ids == self.mask_id
- if not mask_index.any():
- out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
- return out.logits_output, [], out.can_run_graph
+ def step(
+ self,
+ forward_batch: ForwardBatch,
+ full_logits: torch.Tensor,
+ states: List[Any],
+ ) -> List[bool]:
+ batch_size = forward_batch.batch_size
+ done: List[bool] = []
- start_list = []
- prompt_masks = []
for i in range(batch_size):
+ state = states[i]
+ if state["finished"]:
+ done.append(True)
+ continue
+
block_start = i * self.block_size
block_end = block_start + self.block_size
- block_input_ids = forward_batch.input_ids[block_start:block_end]
+ curr_input_ids = forward_batch.input_ids[block_start:block_end]
+ curr_logits = full_logits[block_start:block_end]
+ curr_prompt_mask = state["prompt_mask"]
- prompt_mask = block_input_ids != self.mask_id
- prompt_masks.append(prompt_mask)
- start_list.append(prompt_mask.sum().item())
-
- post_edit_steps = torch.zeros(batch_size, dtype=torch.int32, device=device)
-
- finished = torch.zeros(batch_size, dtype=torch.bool, device=device)
- # Controls whether to perform an additional forward pass for KV cache persistence.
- # For certain decoding rounds where the terminal step yields no state change,
- # this can be set to False to bypass the overhead of an idle forward pass.
- any_changed_in_last_step = False
-
- max_iterations = self.block_size + self.max_post_edit_steps
- for _ in range(max_iterations):
- if finished.all():
- break
-
- out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
- logits_output, can_run_cuda_graph = out.logits_output, out.can_run_graph
-
- any_changed_in_last_step = False
-
- for i in range(batch_size):
- if finished[i]:
- continue
-
- block_start = i * self.block_size
- block_end = block_start + self.block_size
-
- curr_input_ids = forward_batch.input_ids[block_start:block_end]
- curr_logits = logits_output.full_logits[block_start:block_end]
- curr_prompt_mask = prompt_masks[i]
-
- if self.penalty_lambda > 0:
- prev_ids = curr_input_ids[:-1]
- curr_logits[1:, :].scatter_(
- 1, prev_ids.unsqueeze(-1), -self.penalty_lambda, reduce="add"
- )
-
- x = torch.argmax(curr_logits, dim=-1)
- p = torch.squeeze(
- torch.gather(
- F.softmax(curr_logits, dim=-1),
- dim=-1,
- index=torch.unsqueeze(x, -1),
- ),
- -1,
+ if self.penalty_lambda > 0:
+ prev_ids = curr_input_ids[:-1]
+ curr_logits[1:, :].scatter_(
+ 1, prev_ids.unsqueeze(-1), -self.penalty_lambda, reduce="add"
)
- mask_index = curr_input_ids == self.mask_id
- has_mask = mask_index.any()
+ x = torch.argmax(curr_logits, dim=-1)
+ p = torch.squeeze(
+ torch.gather(
+ F.softmax(curr_logits, dim=-1),
+ dim=-1,
+ index=torch.unsqueeze(x, -1),
+ ),
+ -1,
+ )
- # Mask to token (M2T)
- mask_transfer_index = torch.zeros_like(mask_index)
- if has_mask:
- confidence = torch.where(mask_index, p, -np.inf)
- mask_transfer_index = confidence > self.threshold
+ mask_index = curr_input_ids == self.mask_id
+ has_mask = mask_index.any()
- if not mask_transfer_index.any():
- _, select_index = torch.topk(confidence, k=1)
- mask_transfer_index[select_index] = True
- else:
- post_edit_steps[i] += 1
- if post_edit_steps[i] > self.max_post_edit_steps:
- finished[i] = True
- continue
+ # Mask to token (M2T)
+ mask_transfer_index = torch.zeros_like(mask_index)
+ budget_exhausted = False
+ if has_mask:
+ confidence = torch.where(mask_index, p, -np.inf)
+ mask_transfer_index = confidence > self.threshold
+ if not mask_transfer_index.any():
+ _, select_index = torch.topk(confidence, k=1)
+ mask_transfer_index[select_index] = True
+ else:
+ state["post_edit_steps"] += 1
+ if state["post_edit_steps"] > self.max_post_edit_steps:
+ state["finished"] = True
+ budget_exhausted = True
+ if not budget_exhausted:
# Token to token (T2T)
edit_mask = ~mask_index & ~curr_prompt_mask
edit_transfer_index = (
(p > self.edit_threshold) & (curr_input_ids != x) & edit_mask
)
-
transfer_index = mask_transfer_index | edit_transfer_index
- if not transfer_index.any():
- finished[i] = True
- continue
+ if transfer_index.any():
+ curr_input_ids[transfer_index] = x[transfer_index]
+ else:
+ state["finished"] = True
- curr_input_ids[transfer_index] = x[transfer_index]
- any_changed_in_last_step = True
+ # A terminating step changes nothing, so this forward already holds the
+ # block's final KV: emit it now rather than after an extra forward.
+ done.append(state["finished"])
- if any_changed_in_last_step:
- out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
- logits_output, can_run_cuda_graph = out.logits_output, out.can_run_graph
-
- next_token_ids = torch.reshape(forward_batch.input_ids, (batch_size, -1))
- next_token_ids_list = [
- next_token_ids[i, start_list[i] :] for i in range(batch_size)
- ]
-
- return logits_output, next_token_ids_list, can_run_cuda_graph
+ return done
Algorithm = JointThreshold
diff --git a/python/sglang/srt/dllm/algorithm/low_confidence.py b/python/sglang/srt/dllm/algorithm/low_confidence.py
index 6acdcc7f8..418546781 100644
--- a/python/sglang/srt/dllm/algorithm/low_confidence.py
+++ b/python/sglang/srt/dllm/algorithm/low_confidence.py
@@ -1,104 +1,55 @@
-from typing import List, Tuple, Union
+from typing import Any, List
-import numpy as np
import torch
-import torch.nn.functional as F
from sglang.srt.dllm.algorithm.base import DllmAlgorithm
from sglang.srt.dllm.config import DllmConfig
-from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
-from sglang.srt.model_executor.model_runner import ModelRunner
class LowConfidence(DllmAlgorithm):
+ """Each step unmasks positions whose predicted-token confidence exceeds a
+ threshold (falling back to the highest-confidence masked position).
+ """
- def __init__(
- self,
- config: DllmConfig,
- ):
+ def __init__(self, config: DllmConfig):
super().__init__(config)
self.threshold = config.algorithm_config.get("threshold", 0.95)
- def run(
+ def step(
self,
- model_runner: ModelRunner,
forward_batch: ForwardBatch,
- ) -> Tuple[Union[LogitsProcessorOutput, torch.Tensor], List[torch.Tensor], bool]:
+ full_logits: torch.Tensor,
+ states: List[Any],
+ ) -> List[bool]:
batch_size = forward_batch.batch_size
- # Here, the forward_batch full logits contains all the blocks
- # such as [dllm_block_size * batch_size, hidden_size]
- start_list = []
- mask_index = forward_batch.input_ids == self.mask_id
+ vocab_size = full_logits.shape[-1]
+ logits = full_logits.view(batch_size, self.block_size, vocab_size)
+ input_ids = forward_batch.input_ids.view(batch_size, self.block_size)
+ block_mask_index = input_ids == self.mask_id
+ done = block_mask_index.sum(dim=1) == 0
- # Fast path: if there is no mask token, forward and save kv cache
- if torch.sum(mask_index).item() == 0:
- out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
- logits_output, can_run_cuda_graph = out.logits_output, out.can_run_graph
+ x = torch.argmax(logits, dim=-1)
+ probs = torch.nn.functional.softmax(logits, dim=-1)
+ confidence = torch.gather(probs, dim=-1, index=x.unsqueeze(-1)).squeeze(-1)
+ confidence = torch.where(block_mask_index, confidence, -float("inf"))
- next_token_ids = []
- return logits_output, next_token_ids, can_run_cuda_graph
+ transfer_index = confidence > self.threshold
+ has_transfer = transfer_index.sum(dim=1) > 0
+ top1_indices = torch.argmax(confidence, dim=1)
+ batch_indices = torch.arange(batch_size, device=top1_indices.device)
+ top1_mask = torch.zeros_like(transfer_index, dtype=torch.bool)
+ top1_mask[batch_indices, top1_indices] = True
+ transfer_index = torch.where(
+ has_transfer.unsqueeze(-1), transfer_index, top1_mask
+ )
- # Calculate start positions for each block
- for block_id in range(batch_size):
- block_start = block_id * self.block_size
- block_end = block_start + self.block_size
- block_input_ids = forward_batch.input_ids[block_start:block_end]
- block_mask_index = block_input_ids == self.mask_id
- start = self.block_size - torch.sum(block_mask_index).item()
- start_list.append(start)
+ x = torch.where(block_mask_index, x, input_ids)
+ new_input_ids = torch.where(transfer_index, x, input_ids)
+ # In-place to preserve the input_ids tensor identity (CUDA graph safe).
+ forward_batch.input_ids.copy_(new_input_ids.view(-1))
- for _ in range(self.block_size):
- mask_index = forward_batch.input_ids == self.mask_id
- if torch.sum(mask_index).item() == 0:
- break
-
- out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
- logits_output, can_run_cuda_graph = out.logits_output, out.can_run_graph
- assert batch_size == forward_batch.input_ids.shape[0] // self.block_size
- for batch_id in range(batch_size):
- curr_block_start = batch_id * self.block_size
- curr_block_end = curr_block_start + self.block_size
- block_input_ids = forward_batch.input_ids[
- curr_block_start:curr_block_end,
- ]
- block_mask_index = block_input_ids == self.mask_id
- if torch.sum(block_mask_index).item() == 0:
- continue
- curr_logits = logits_output.full_logits[
- curr_block_start:curr_block_end,
- ]
-
- x = torch.argmax(curr_logits, dim=-1)
- p = torch.squeeze(
- torch.gather(
- F.softmax(curr_logits, dim=-1),
- dim=-1,
- index=torch.unsqueeze(x, -1),
- ),
- -1,
- )
- x = torch.where(block_mask_index, x, block_input_ids)
- confidence = torch.where(block_mask_index, p, -np.inf)
-
- transfer_index = confidence > self.threshold
-
- if transfer_index.sum().item() == 0:
- _, select_index = torch.topk(confidence, k=1)
- transfer_index[select_index] = True
-
- block_input_ids[transfer_index] = x[transfer_index]
-
- out = model_runner.forward(forward_batch, pp_proxy_tensors=None)
- logits_output, can_run_cuda_graph = out.logits_output, out.can_run_graph
- # Here next token ids is tricky to implement the dynamic lengths,
- # so we return a list of tensors
- next_token_ids = torch.reshape(forward_batch.input_ids, (batch_size, -1))
- next_token_ids_list = [
- next_token_ids[i, start_list[i] :] for i in range(batch_size)
- ]
-
- return logits_output, next_token_ids_list, can_run_cuda_graph
+ return done.tolist()
Algorithm = LowConfidence
diff --git a/python/sglang/srt/dllm/config.py b/python/sglang/srt/dllm/config.py
index edd204926..f0f2b9d11 100644
--- a/python/sglang/srt/dllm/config.py
+++ b/python/sglang/srt/dllm/config.py
@@ -12,12 +12,14 @@ class DllmConfig:
block_size: int,
mask_id: int,
max_running_requests: int,
+ first_done_first_out_mode: bool = False,
):
self.algorithm = algorithm
self.algorithm_config = algorithm_config
self.block_size = block_size
self.mask_id = mask_id
self.max_running_requests = max_running_requests
+ self.first_done_first_out_mode = first_done_first_out_mode
@staticmethod
def from_server_args(
@@ -72,4 +74,5 @@ class DllmConfig:
block_size=block_size,
mask_id=mask_id,
max_running_requests=max_running_requests,
+ first_done_first_out_mode=server_args.dllm_fdfo,
)
diff --git a/python/sglang/srt/dllm/mixin/req.py b/python/sglang/srt/dllm/mixin/req.py
index 74635d1c4..ff72cd4fc 100644
--- a/python/sglang/srt/dllm/mixin/req.py
+++ b/python/sglang/srt/dllm/mixin/req.py
@@ -20,6 +20,8 @@ class DllmReqPhase(str, enum.Enum):
class ReqDllmMixin:
def init_diffusion_llm(self: Req, dllm_config: DllmConfig):
self.dllm_phase: Optional[DllmReqPhase] = None
+ self.dllm_incomplete_ids = array("q")
+ self.dllm_algo_state = None
self.dllm_block_offset = 0
self.dllm_config = dllm_config
@@ -39,6 +41,10 @@ class ReqDllmMixin:
]
def determine_dllm_phase(self: Req):
+ if self.dllm_incomplete_ids:
+ self.dllm_phase = DllmReqPhase.STAGING_DECODE
+ return
+
prefix_length = len(self.prefix_indices)
min_required_length = prefix_length + self.dllm_config.block_size
@@ -55,6 +61,17 @@ class ReqDllmMixin:
self.dllm_phase = DllmReqPhase.STAGING_DECODE
def _init_fill_ids_for_dllm(self: Req):
+ if self.dllm_incomplete_ids:
+ prefix_len = len(self.prefix_indices)
+ assert len(self.dllm_incomplete_ids) == self.dllm_config.block_size
+ self.full_untruncated_fill_ids = (
+ self.full_untruncated_fill_ids[:prefix_len] + self.dllm_incomplete_ids
+ )
+ # extend_range is (re)computed by the staging adder
+ # (add_dllm_staging_req) before this req is scheduled, mirroring the
+ # non-incomplete path which also defers it to the adder.
+ return
+
self.dllm_block_offset = (
0
if not self.dllm_initialized
diff --git a/python/sglang/srt/dllm/mixin/scheduler.py b/python/sglang/srt/dllm/mixin/scheduler.py
index b41ff1ccf..f9597a02d 100644
--- a/python/sglang/srt/dllm/mixin/scheduler.py
+++ b/python/sglang/srt/dllm/mixin/scheduler.py
@@ -72,24 +72,84 @@ class SchedulerDllmMixin:
if result.copy_done is not None:
result.copy_done.synchronize()
- if result.next_token_ids:
- self.token_to_kv_pool_allocator.free_group_begin()
+ fdfo_mode = self.dllm_config.first_done_first_out_mode
+ assert (
+ not fdfo_mode or result.accept_length_per_req_cpu is not None
+ ), "FDFO dLLM result is missing accept lengths."
+ # Sync mode emits tokens only once a block fully resolves; FDFO always
+ # commits (resolved blocks decode, unresolved blocks stash + free KV).
+ if fdfo_mode or result.next_token_ids:
+ block_size = self.dllm_config.block_size
+ algo_states = result.dllm_algo_state
+
+ self.token_to_kv_pool_allocator.free_group_begin()
for idx in range(batch.batch_size()):
req = batch.reqs[idx]
- next_token_ids = result.next_token_ids[idx].tolist()
- new_tokens = len(next_token_ids)
- if new_tokens == 0:
+ if not fdfo_mode:
+ next_token_ids = result.next_token_ids[idx].tolist()
+ new_tokens = len(next_token_ids)
+ if new_tokens == 0:
+ continue
+
+ req.full_untruncated_fill_ids[
+ req.extend_range.end - new_tokens : req.extend_range.end
+ ] = array("q", next_token_ids)
+ self.metrics_reporter.num_generated_tokens += new_tokens
+
+ req.output_ids.extend(next_token_ids)
+ req.update_finish_state(new_accepted_len=new_tokens)
+
+ if req.finished():
+ release_kv_cache(req, self.tree_cache)
+ req.time_stats.set_completion_time()
continue
- req.full_untruncated_fill_ids[
- req.extend_range.end - new_tokens : req.extend_range.end
- ] = array("q", next_token_ids)
- self.metrics_reporter.num_generated_tokens += new_tokens
+ next_token_ids = result.next_token_ids[idx]
+ assert len(next_token_ids) == block_size
+ if result.accept_length_per_req_cpu[idx] == 0:
+ # Block unresolved: stash partial state and free the KV slots
+ # of the still-masked block so the next FDFO round can
+ # re-denoise it without leaking the previous allocation.
+ req.dllm_incomplete_ids = array("q", next_token_ids)
+ req.dllm_algo_state = (
+ algo_states[idx] if algo_states is not None else None
+ )
+ old_prefix_len = len(req.prefix_indices)
+ new_fill_len = req.extend_range.end
+ if new_fill_len > old_prefix_len:
+ kv_indices_to_free = self.req_to_token_pool.req_to_token[
+ req.req_pool_idx, old_prefix_len:new_fill_len
+ ]
+ self.token_to_kv_pool_allocator.free(kv_indices_to_free)
+ continue
+
+ req.dllm_incomplete_ids = array("q")
+ req.dllm_algo_state = None
+
+ # Mirror the resolved block into the committed fill ids so the
+ # prefix cache keys on the real tokens, not the mask block, next
+ # round. Index relative to extend_range.end (the truncated/
+ # committed length), which can be shorter than
+ # full_untruncated_fill_ids when the staging adder truncates the
+ # block to the KV budget.
+ req.full_untruncated_fill_ids[
+ req.extend_range.end - block_size : req.extend_range.end
+ ] = array("q", next_token_ids)
+
+ len_input = len(req.origin_input_ids)
+ len_fill = req.extend_range.end
+ if len_fill <= len_input:
+ continue
+
+ if len_fill - len(next_token_ids) < len_input:
+ next_token_ids = next_token_ids[len_input - len_fill :]
+
+ self.metrics_reporter.num_generated_tokens += len(next_token_ids)
req.output_ids.extend(next_token_ids)
- req.update_finish_state(new_accepted_len=new_tokens)
+ req.update_finish_state(new_accepted_len=len(next_token_ids))
if req.finished():
release_kv_cache(req, self.tree_cache)
@@ -98,11 +158,10 @@ class SchedulerDllmMixin:
self.output_streamer.stream_output(batch.reqs, batch.return_logprob)
self.token_to_kv_pool_allocator.free_group_end()
- can_run_cuda_graph = result.can_run_cuda_graph
self.metrics_reporter.report_prefill_stats(
batch=batch,
prefill_stats=batch.prefill_stats,
- can_run_cuda_graph=can_run_cuda_graph,
+ can_run_cuda_graph=result.can_run_cuda_graph,
dp_cooperation_info=batch.dp_cooperation_info,
)
diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py
index d464e74af..00f97dcaa 100644
--- a/python/sglang/srt/managers/scheduler.py
+++ b/python/sglang/srt/managers/scheduler.py
@@ -2624,7 +2624,12 @@ class Scheduler(
if self.dllm_config is not None and self.dllm_manager.any_staging_reqs():
chunked_req_to_exclude.update(self.dllm_manager.staging_queue)
for req in self.dllm_manager.staging_queue:
- self.stash_chunked_request(req)
+ if self.dllm_config.first_done_first_out_mode:
+ if not req.dllm_incomplete_ids:
+ self.stash_chunked_request(req)
+ self.req_to_token_pool.free(req)
+ else:
+ self.stash_chunked_request(req)
if self.chunked_req is not None:
# Move the chunked request out of the batch so that we can merge
diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py
index 3db49daf8..65d8e5ae3 100644
--- a/python/sglang/srt/managers/tp_worker.py
+++ b/python/sglang/srt/managers/tp_worker.py
@@ -475,14 +475,27 @@ class TpModelWorker(BaseTpWorker):
return self.dllm_algorithm is not None
def _forward_batch_generation_dllm(
- self, forward_batch: ForwardBatch
+ self,
+ forward_batch: ForwardBatch,
+ batch: Optional[ScheduleBatch] = None,
) -> GenerationBatchResult:
- logits_output, next_token_ids, can_run_cuda_graph = self.dllm_algorithm.run(
- self.model_runner, forward_batch
- )
+ algo_states = None
+ if self.dllm_algorithm.fdfo and batch is not None:
+ algo_states = [req.dllm_algo_state for req in batch.reqs]
+
+ (
+ logits_output,
+ next_token_ids,
+ accept_length_per_req_cpu,
+ dllm_algo_state,
+ can_run_cuda_graph,
+ ) = self.dllm_algorithm.run(self.model_runner, forward_batch, algo_states)
+
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=next_token_ids,
+ accept_length_per_req_cpu=accept_length_per_req_cpu,
+ dllm_algo_state=dllm_algo_state,
can_run_cuda_graph=can_run_cuda_graph,
)
@@ -508,7 +521,7 @@ class TpModelWorker(BaseTpWorker):
forward_batch.apply_deprecated_skip_attn_backend_init(skip_attn_backend_init)
if self.is_dllm():
- return self._forward_batch_generation_dllm(forward_batch)
+ return self._forward_batch_generation_dllm(forward_batch, batch)
if self.pp_group.is_last_rank:
out = self.model_runner.forward(
diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py
index 87fbb35de..c932de486 100644
--- a/python/sglang/srt/managers/utils.py
+++ b/python/sglang/srt/managers/utils.py
@@ -39,9 +39,14 @@ def _async_d2h(t: torch.Tensor) -> torch.Tensor:
class GenerationBatchResult:
logits_output: Optional[LogitsProcessorOutput] = None
pp_hidden_states_proxy_tensors: Optional[PPProxyTensors] = None
- next_token_ids: Optional[Union[torch.Tensor, List[torch.Tensor]]] = None
+ next_token_ids: Optional[
+ Union[torch.Tensor, List[torch.Tensor], List[List[int]]]
+ ] = None
num_correct_drafts: int = 0 # no bonus included
num_correct_drafts_per_req_cpu: Optional[List[int]] = None
+ # FDFO dLLM batching: per-request accepted block length and carried algo state.
+ accept_length_per_req_cpu: Optional[List[int]] = None
+ dllm_algo_state: Optional[List[Any]] = None
can_run_cuda_graph: bool = False
# PP skip output comm: True when output send/recv was skipped and
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index 3461b8aeb..0ad76225d 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -2349,6 +2349,13 @@ class ServerArgs:
Optional[str],
"The diffusion LLM algorithm configurations. Must be a YAML file.",
] = None
+ dllm_fdfo: A[
+ bool,
+ Arg(
+ help="Enable First-Done-First-Out (FDFO) scheduling for diffusion LLM inference. Enabled by default; use --no-dllm-fdfo to fall back to synchronous block scheduling.",
+ action=argparse.BooleanOptionalAction,
+ ),
+ ] = True
# -------------------------------------------------------------------------
# PD disaggregation
diff --git a/test/registered/dllm/test_llada2_mini.py b/test/registered/dllm/test_dllm_batching_fdfo.py
similarity index 89%
rename from test/registered/dllm/test_llada2_mini.py
rename to test/registered/dllm/test_dllm_batching_fdfo.py
index 7b869c70c..5607d91be 100644
--- a/test/registered/dllm/test_llada2_mini.py
+++ b/test/registered/dllm/test_dllm_batching_fdfo.py
@@ -20,7 +20,9 @@ from sglang.test.test_utils import (
)
-class TestLLaDA2Mini(CustomTestCase):
+class TestBatchingFDFO(CustomTestCase):
+ """End-to-end dLLM coverage on the default First-Done-First-Out scheduler."""
+
@classmethod
def setUpClass(cls):
cls.model = "inclusionAI/LLaDA2.0-mini"
@@ -38,6 +40,7 @@ class TestLLaDA2Mini(CustomTestCase):
"flashinfer",
"--dllm-algorithm",
"LowConfidence",
+ "--dllm-fdfo",
"--cuda-graph-bs",
"1",
"2",
@@ -73,7 +76,7 @@ class TestLLaDA2Mini(CustomTestCase):
if is_in_amd_ci():
self.assertGreater(metrics["output_throughput"], 80)
else:
- self.assertGreater(metrics["output_throughput"], 350)
+ self.assertGreater(metrics["output_throughput"], 450)
def test_bs_1_speed(self):
args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048)
@@ -83,7 +86,7 @@ class TestLLaDA2Mini(CustomTestCase):
if is_in_ci():
write_github_step_summary(
- f"### test_bs_1_speed (llada2-mini) with tp1\n"
+ f"### test_bs_1_speed (llada2-mini FDFO) with tp1\n"
f"{speed=:.2f} token/s\n"
)
if is_in_amd_ci():
diff --git a/test/registered/dllm/test_dllm_fdfo_joint_threshold.py b/test/registered/dllm/test_dllm_fdfo_joint_threshold.py
new file mode 100644
index 000000000..a5671d93f
--- /dev/null
+++ b/test/registered/dllm/test_dllm_fdfo_joint_threshold.py
@@ -0,0 +1,89 @@
+from sglang.test.ci.ci_register import register_cuda_ci
+
+register_cuda_ci(est_time=300, stage="base-b", runner_config="1-gpu-large")
+
+import unittest
+
+import requests
+
+from sglang.srt.utils import kill_process_tree
+from sglang.test.test_utils import (
+ DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
+ DEFAULT_URL_FOR_TEST,
+ CustomTestCase,
+ popen_launch_server,
+)
+
+PROMPTS = [
+ "Question: Natalia sold clips to 48 friends in April, and half as many in "
+ "May. How many clips did she sell altogether? Answer:",
+ "The capital of France is",
+ "Q: What is 12 times 13? A:",
+]
+
+
+class TestBatchingFDFOJointThreshold(CustomTestCase):
+ """At a single in-flight request, FDFO and synchronous execution run identical
+ forward shapes, so a correct stateful (``dllm_algo_state``) carry must produce
+ byte-identical multi-block output.
+ """
+
+ model = "inclusionAI/LLaDA2.1-mini"
+ base_url = DEFAULT_URL_FOR_TEST
+
+ def _collect_outputs(self, fdfo: bool):
+ other_args = [
+ "--trust-remote-code",
+ "--tp-size",
+ "1",
+ "--mem-fraction-static",
+ "0.9",
+ "--max-running-requests",
+ "1",
+ "--attention-backend",
+ "flashinfer",
+ "--dllm-algorithm",
+ "JointThreshold",
+ "--cuda-graph-bs",
+ "1",
+ ]
+ # FDFO is the default; the sync arm must opt out explicitly.
+ other_args.append("--dllm-fdfo" if fdfo else "--no-dllm-fdfo")
+
+ process = popen_launch_server(
+ self.model,
+ self.base_url,
+ timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
+ other_args=other_args,
+ )
+ try:
+ outputs = []
+ for prompt in PROMPTS:
+ response = requests.post(
+ f"{self.base_url}/v1/completions",
+ json={
+ "model": self.model,
+ "prompt": prompt,
+ "max_tokens": 128,
+ "temperature": 0,
+ },
+ timeout=120,
+ )
+ outputs.append(response.json()["choices"][0]["text"])
+ return outputs
+ finally:
+ kill_process_tree(process.pid)
+
+ def test_fdfo_matches_sync(self):
+ sync_outputs = self._collect_outputs(fdfo=False)
+ fdfo_outputs = self._collect_outputs(fdfo=True)
+ self.assertEqual(
+ fdfo_outputs,
+ sync_outputs,
+ "JointThreshold FDFO output must match synchronous output, which "
+ "validates the cross-step dllm_algo_state carry across blocks.",
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()