Skip to content

vllm.model_executor.models.diffusion_gemma

DiffusionGemma model, ModelState, and Sampler for vLLM.

Single Gemma4 backbone run in two modes (like YOCO): - encoder mode: causal attention, writes KV cache - decoder mode: bidirectional attention, reads encoder KV, doesn't write

Same weights, same layers. The only decoder-unique component is a self-conditioning MLP.

Multimodal support: the model always includes a vision tower (shared with Gemma4). Images are encoded through the vision tower and projected into the LM embedding space via Gemma4MultimodalEmbedder.

Classes:

DiffusionGemmaForConditionalGeneration

Bases: Module, SupportsMultiModal, SupportsQuant

DiffusionGemma for vLLM.

Single Gemma4 backbone that switches between encoder and decoder mode. The encoder path uses standard Gemma4 layers (causal attention, KV write). The decoder path uses the same weights with bidirectional attention and KV read-only, plus self-conditioning.

Always includes a vision tower (same as Gemma4) for image understanding.

In practice, the model's forward() dispatches based on the mode kwarg set by DiffusionGemmaModelState.prepare_inputs().

Methods:

  • get_mm_mapping –

    Get the module prefix mapping for multimodal models.

Source code in vllm/model_executor/models/diffusion_gemma.py
@MULTIMODAL_REGISTRY.register_processor(
    Gemma4MultiModalProcessor,
    info=DiffusionGemmaProcessingInfo,
    dummy_inputs=Gemma4DummyInputsBuilder,
)
class DiffusionGemmaForConditionalGeneration(
    nn.Module,
    SupportsMultiModal,
    SupportsQuant,
):
    """DiffusionGemma for vLLM.

    Single Gemma4 backbone that switches between encoder and decoder mode.
    The encoder path uses standard Gemma4 layers (causal attention, KV write).
    The decoder path uses the same weights with bidirectional attention and
    KV read-only, plus self-conditioning.

    Always includes a vision tower (same as Gemma4) for image understanding.

    In practice, the model's forward() dispatches based on the `mode` kwarg
    set by DiffusionGemmaModelState.prepare_inputs().
    """

    hf_to_vllm_mapper = WeightsMapper(
        orig_to_new_prefix={
            "model.decoder.self_conditioning.": "self_conditioning.",
            "model.decoder.": "model.",
            "model.encoder.language_model.": "model.",
            "model.encoder.vision_tower.": "vision_tower.",
            "model.encoder.embed_vision.": "embed_vision.",
        },
    )

    packed_modules_mapping = {
        "qkv_proj": ["q_proj", "k_proj", "v_proj"],
        "gate_up_proj": ["gate_proj", "up_proj"],
    }

    @staticmethod
    def get_model_state_cls():
        return DiffusionGemmaModelState

    def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
        super().__init__()
        config = vllm_config.model_config.hf_config
        text_config = vllm_config.model_config.hf_text_config
        self.config = config
        self.model_dtype = vllm_config.model_config.dtype

        # DiffusionGemma's full-attention layers have NO v_proj — V is
        # computed from k_proj's output (`value_states = key_states` before
        # k_norm in `DiffusionGemmaDecoderTextAttention.forward`). This is
        # the "k_eq_v" variant in our Gemma4 backbone. The checkpoint has no
        # v_proj weights for full-attention layers; without this flag they
        # would silently load with random V projections.
        text_config.attention_k_eq_v = True

        # ---- Vision tower ----
        # Gemma4's image path, borrowed below, reads this flag.
        lora_config = vllm_config.lora_config
        self._enable_mm_lora = bool(
            lora_config is not None and lora_config.enable_tower_connector_lora
        )
        vision_config = getattr(config, "vision_config", None)
        self.embed_vision: Gemma4MultimodalEmbedder | None
        if vision_config is not None:
            quant_config = vllm_config.quant_config
            tower_quant: QuantizationConfig | None
            if quant_config and quant_config.get_name() in [
                "bitsandbytes",
                "torchao",
                "compressed-tensors",
            ]:
                tower_quant = quant_config
            else:
                quantizable = (
                    vision_config.hidden_size % 64 == 0
                    and vision_config.intermediate_size % 64 == 0
                )
                tower_quant = quant_config if quantizable else None

            with self._mark_tower_model(vllm_config, {"image", "video"}):
                self.vision_tower = AutoModel.from_config(config=vision_config)
                self.embed_vision = Gemma4MultimodalEmbedder(
                    vision_config,
                    text_config,
                    quant_config=tower_quant,
                    prefix=maybe_prefix(prefix, "embed_vision"),
                )
                recursive_replace_linear(
                    self.vision_tower,
                    tower_quant,
                    prefix=maybe_prefix(prefix, "vision_tower"),
                )
        else:
            self.vision_tower = None
            self.embed_vision = None

        # ---- Language backbone (Gemma4Model) ----
        # Use maybe_prefix to ensure correct weight name prefixes for
        # quantization. The quantization config uses hf_to_vllm_mapper to
        # match checkpoint weight names to model parameter names.
        self.model = Gemma4Model(
            vllm_config=vllm_config,
            prefix=maybe_prefix(prefix, "model"),
        )

        self.lm_head = ParallelLMHead(
            num_embeddings=text_config.vocab_size,
            embedding_dim=text_config.hidden_size,
            quant_config=vllm_config.quant_config,
            prefix=maybe_prefix(prefix, "lm_head"),
        )

        if text_config.tie_word_embeddings:
            self.lm_head = self.lm_head.tie_weights(self.model.embed_tokens)

        # HF DiffusionGemma applies the final-logit softcap in fp32, before
        # any other processing. Do it manually in `compute_logits` so the
        # LogitsProcessor only handles the lm_head GEMM.
        self.final_logit_softcapping = getattr(
            text_config, "final_logit_softcapping", None
        )
        self.logits_processor = LogitsProcessor(
            text_config.vocab_size,
            soft_cap=None,
        )

        sc_size = (
            getattr(config, "self_conditioning_size", None)
            or text_config.intermediate_size
        )
        self.self_conditioning = DiffusionGemmaSelfConditioning(
            hidden_size=text_config.hidden_size,
            self_conditioning_size=sc_size,
            eps=getattr(text_config, "rms_norm_eps", 1e-6),
        )

    def compute_self_conditioning(
        self,
        inputs_embeds: torch.Tensor,
        probs: torch.Tensor,
    ) -> torch.Tensor:
        embed_weight = self.model.embed_tokens.weight
        soft_embeds = torch.matmul(
            probs.to(embed_weight.dtype), embed_weight
        ) * self.model.normalizer.to(inputs_embeds.dtype)
        return self.self_conditioning(inputs_embeds, soft_embeds)

    # ------------------------------------------------------------------ #
    # Multimodal: reuse Gemma4's image parsing, processing & embedding
    # ------------------------------------------------------------------ #
    # The vision tower, pooler, embed_vision, and their processing logic
    # are architecturally identical to Gemma4.  Delegate to avoid
    # maintaining a duplicate copy.

    _parse_and_validate_image_input = (
        Gemma4ForConditionalGeneration._parse_and_validate_image_input
    )
    _parse_and_validate_video_input = (
        Gemma4ForConditionalGeneration._parse_and_validate_video_input
    )
    _parse_and_validate_multimodal_inputs = (
        Gemma4ForConditionalGeneration._parse_and_validate_multimodal_inputs
    )
    _encoder_chunk = staticmethod(Gemma4ForConditionalGeneration._encoder_chunk)
    _process_image_input = Gemma4ForConditionalGeneration._process_image_input
    _process_video_input = Gemma4ForConditionalGeneration._process_video_input
    embed_multimodal = Gemma4ForConditionalGeneration.embed_multimodal

    def get_mm_mapping(self) -> MultiModelKeys:
        """Get the module prefix mapping for multimodal models."""
        return MultiModelKeys.from_string_field(
            language_model="model",
            connector=["embed_vision"],
            tower_model=["vision_tower"],
        )

    # ------------------------------------------------------------------ #
    # Forward
    # ------------------------------------------------------------------ #

    def forward(
        self,
        input_ids: torch.Tensor,
        positions: torch.Tensor,
        intermediate_tensors: Any | None = None,
        inputs_embeds: torch.Tensor | None = None,
        **kwargs: Any,
    ) -> torch.Tensor:
        if intermediate_tensors is not None:
            inputs_embeds = None
        return self.model(
            input_ids=input_ids,
            positions=positions,
            intermediate_tensors=intermediate_tensors,
            inputs_embeds=inputs_embeds,
            **kwargs,
        )

    def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor | None:
        states = getattr(self, "diffusion_states", None)
        allowed = states.step_allowed if states is not None else None
        if allowed is not None:
            # One column per allowed id: a [rows, K] GEMM over the K gathered
            # rows of the tied embedding.
            logits = torch.nn.functional.linear(
                hidden_states, self.lm_head.weight[allowed]
            ).float()
            if self.final_logit_softcapping is not None:
                logits = _softcap_logits(logits, self.final_logit_softcapping)
            return logits
        logits = self.logits_processor(self.lm_head, hidden_states)
        if logits is not None and self.final_logit_softcapping is not None:
            logits = _softcap_logits(logits, self.final_logit_softcapping)
        return logits

    def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
        # Some checkpoints carry a vestigial Gemma3n-style embedding table.
        loader = AutoWeightsLoader(
            self, ignore_unexpected_prefixes=["embed_vision.embedding."]
        )
        return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper)

    @classmethod
    def get_placeholder_str(cls, modality: str, i: int) -> str | None:
        if modality == "image":
            return "<image_soft_token>"
        if modality == "video":
            return "<|video|>"
        raise ValueError(f"Unsupported modality: {modality}")

get_mm_mapping()

Get the module prefix mapping for multimodal models.

Source code in vllm/model_executor/models/diffusion_gemma.py
def get_mm_mapping(self) -> MultiModelKeys:
    """Get the module prefix mapping for multimodal models."""
    return MultiModelKeys.from_string_field(
        language_model="model",
        connector=["embed_vision"],
        tower_model=["vision_tower"],
    )

DiffusionGemmaModelState

Bases: ModelState

ModelState for DiffusionGemma.

Single Gemma4 backbone in two modes: - encoder mode (num_draft_tokens == 0): causal attention, writes KV - decoder mode (num_draft_tokens > 0): bidirectional attention, reads KV

Source code in vllm/model_executor/models/diffusion_gemma.py
class DiffusionGemmaModelState(ModelState):
    """ModelState for DiffusionGemma.

    Single Gemma4 backbone in two modes:
    - encoder mode (num_draft_tokens == 0): causal attention, writes KV
    - decoder mode (num_draft_tokens > 0): bidirectional attention, reads KV
    """

    def __init__(
        self,
        vllm_config: VllmConfig,
        model: nn.Module,
        encoder_cache: Any,
        device: torch.device,
    ) -> None:
        super().__init__(vllm_config, model, encoder_cache, device)

        # Per-step MM data produced by prepare_inputs_embeds and consumed by
        # prepare_inputs.  Stored as raw (mm_embeds, is_mm_embed) so that
        # prepare_inputs can call embed_input_ids directly into the
        # persistent _inputs_embeds_buf, avoiding the intermediate copy
        # through encoder_runner.inputs_embeds.
        self._pending_mm_embeds: tuple[list[torch.Tensor], torch.Tensor] | None = None

        diffusion_config = vllm_config.diffusion_config
        canvas_length = diffusion_config.canvas_length if diffusion_config else 32

        text_config = self.model_config.hf_text_config
        self.gen_config = self.model_config.try_get_generation_config()
        max_denoising_steps = (
            diffusion_config.max_denoising_steps if diffusion_config else None
        ) or self.gen_config.get("max_denoising_steps", 48)
        self.diffusion_states = DiffusionGemmaRequestStates(
            max_num_reqs=self.max_num_reqs,
            canvas_length=canvas_length,
            vocab_size=self.model_config.get_vocab_size(),
            max_denoising_steps=max_denoising_steps,
            device=device,
            hidden_size=text_config.hidden_size,
            # In Transformers, `stability_threshold=1` (the default) means the current
            # step must match the previous step. In vLLM, the history buffer includes
            # the current step, so we add 1 to match the same behavior.
            stability_threshold=self.gen_config["stability_threshold"] + 1,
        )
        # compute_logits reads the step's shared allowed set from here.
        self.model.diffusion_states = self.diffusion_states
        self._req_id_to_index: dict[str, int] = {}

        # Persistent buffer for per-request causal flags, updated in-place
        # so FULL CUDA graph replay sees the latest values. Must be int32:
        # FlashAttentionMetadataBuilder.build() rejects other dtypes, since
        # an out-of-place cast there would detach the captured graph from
        # this buffer and freeze replay at the capture-time snapshot.
        self._causal_buf = torch.zeros(
            self.max_num_reqs, dtype=torch.int32, device=device
        )

        # Persistent inputs_embeds buffer — required so FULL CUDA graph
        # capture and runtime point at the SAME memory address.
        # `prepare_dummy_inputs` (capture path) and `prepare_inputs` (runtime
        # path) both must hand the captured graph a tensor at this address.
        self._inputs_embeds_buf = torch.zeros(
            self.max_num_tokens,
            text_config.hidden_size,
            dtype=self.model_config.dtype,
            device=device,
        )

    def get_supported_generation_tasks(self):
        return ("generate",)

    def custom_sampler(self, sampler: Any) -> tuple[Any, Any] | None:
        diffusion_config = self.vllm_config.diffusion_config
        gen = self.gen_config
        sampler_cfg = gen.get("sampler_config") or {}
        if "EntropyBound" not in sampler_cfg.get("_cls_name", ""):
            raise ValueError("DiffusionGemma requires an EntropyBound sampler_config")
        entropy_bound = sampler_cfg.get("entropy_bound")
        if entropy_bound is None or entropy_bound <= 0:
            raise ValueError(
                f"entropy_bound must be a positive float (got {entropy_bound})"
            )
        # The self-conditioning matmul (probs @ embed_tokens.weight) runs over a
        # vocab-parallel embedding shard. Hand the sampler this rank's vocab
        # slice and TP group so it can all-reduce the partial products.
        embed_tokens = self.model.model.embed_tokens
        shard = embed_tokens.shard_indices
        tp_group = get_tp_group()
        return DiffusionSampler(
            sampler=sampler,
            diffusion_config=diffusion_config,
            vocab_size=self.model_config.get_vocab_size(),
            diffusion_states=self.diffusion_states,
            t_min=gen["t_min"],
            t_max=gen["t_max"],
            entropy_bound=entropy_bound,
            confidence_threshold=gen["confidence_threshold"],
            embed_weight=embed_tokens.weight,
            normalizer=self.model.model.normalizer,
            sc_vocab_start=shard.org_vocab_start_index,
            sc_vocab_end=shard.org_vocab_end_index,
            tp_size=tp_group.world_size,
            tp_group_name=tp_group.unique_name,
        ), None

    def apply_staged_writes(self) -> None:
        pass

    def add_request(self, req_index: int, new_req_data: Any) -> None:
        self._req_id_to_index[new_req_data.req_id] = req_index
        self.diffusion_states.add_request(req_index)
        if not new_req_data.req_id.startswith("_warmup_"):
            prompt_len = len(new_req_data.prompt_token_ids)
            self.diffusion_states.prompt_len[req_index].fill_(prompt_len)

    def remove_request(self, req_id: str) -> None:
        idx = self._req_id_to_index.pop(req_id, None)
        if idx is not None:
            self.diffusion_states.remove_request(idx)

    def prepare_inputs_embeds(
        self,
        scheduled_encoder_inputs: dict[str, list[int]],
        input_batch: InputBatch,
        req_states: RequestState,
    ) -> torch.Tensor | None:
        if not self.supports_mm_inputs:
            return None

        mm_hashes, mm_kwargs = self.encoder_runner.prepare_mm_inputs(
            scheduled_encoder_inputs
        )
        if mm_kwargs:
            encoder_outputs = self.encoder_runner.execute_mm_encoder(mm_kwargs)
            self.encoder_cache.encoder_outputs.update(zip(mm_hashes, encoder_outputs))

        mm_embeds, is_mm_embed = self.gather_mm_embeddings(input_batch)

        if not mm_embeds:
            # No MM tokens in this batch (e.g. all-decode step).
            # prepare_inputs will use embed_input_ids (text-only) directly.
            self._pending_mm_embeds = None
            return None

        # Stash raw MM ingredients for prepare_inputs to merge directly
        # into the persistent buffer, avoiding the intermediate copy
        # through encoder_runner.inputs_embeds.
        self._pending_mm_embeds = (mm_embeds, is_mm_embed)
        return None

    def _apply_self_conditioning(
        self,
        decode_slots_np: np.ndarray,
        decode_idx_np: np.ndarray,
        query_start_loc_np: np.ndarray,
        inputs_embeds: torch.Tensor,
        sc_embeds: torch.Tensor,
    ) -> None:
        # One self-conditioning MLP call per decode request, over that request's
        # query span [start, end) = its canvas. The span is the full canvas (CL)
        # or, for the final canvas truncated near max_model_len, fewer than CL
        # positions. sc_embeds already holds probs @ embed_weight from the prior
        # denoise step, masked to zero by the sampler for slots not denoising
        # this step; only the MLP runs here. CPU metadata -> no GPU syncs.
        for slot, idx in zip(decode_slots_np.tolist(), decode_idx_np.tolist()):
            start = int(query_start_loc_np[idx])
            end = int(query_start_loc_np[idx + 1])
            canvas = slice(start, end)
            soft = sc_embeds[slot, : end - start]
            inputs_embeds[canvas] = self.model.self_conditioning(
                inputs_embeds[canvas], soft.to(inputs_embeds.dtype)
            )

    def prepare_inputs(self, input_batch, req_states) -> dict[str, Any]:
        states = self.diffusion_states
        num_tokens = input_batch.num_tokens
        num_reqs = input_batch.num_reqs

        # Write into the PERSISTENT inputs_embeds buffer so FULL CUDA graph
        # replay sees the latest values at the captured address.
        num_tokens_padded = input_batch.num_tokens_after_padding
        inputs_embeds = self._inputs_embeds_buf[:num_tokens_padded]

        # Populate embeddings: merge MM features when available,
        # otherwise embed input_ids as text-only.
        input_ids = input_batch.input_ids[:num_tokens]
        if self._pending_mm_embeds is not None:
            mm_embeds, is_mm_embed = self._pending_mm_embeds
            self._pending_mm_embeds = None
            inputs_embeds[:num_tokens].copy_(
                self.model.embed_input_ids(
                    input_ids,
                    multimodal_embeddings=mm_embeds,
                    is_multimodal=is_mm_embed,
                )
            )
        else:
            inputs_embeds[:num_tokens].copy_(self.model.embed_input_ids(input_ids))

        # Apply self-conditioning ONLY for denoising decode requests.
        states.step_allowed = None
        if input_batch.num_draft_tokens > 0 and self._req_id_to_index:
            slots_np = input_batch.idx_mapping_np[:num_reqs]
            num_logits_np = np.diff(input_batch.cu_num_logits_np[: num_reqs + 1])
            is_decode_indices_np = np.where(num_logits_np > 0)[0]
            states.step_allowed = states.batch_allowed(
                slots_np[is_decode_indices_np].tolist()
            )
            self._apply_self_conditioning(
                slots_np[is_decode_indices_np],
                is_decode_indices_np,
                input_batch.query_start_loc_np,
                inputs_embeds,
                states.self_conditioning_embeds,
            )

        return {"inputs_embeds": inputs_embeds}

    def prepare_dummy_inputs(self, num_reqs: int, num_tokens: int) -> dict[str, Any]:
        # CUDA graph capture path — return a slice of the SAME persistent
        # inputs_embeds buffer that `prepare_inputs` writes to at runtime,
        # so the captured graph and runtime point to identical addresses.
        return {"inputs_embeds": self._inputs_embeds_buf[:num_tokens]}

    def postprocess_state(
        self, idx_mapping, num_sampled, num_computed_tokens=None
    ) -> None:
        return None

    def prepare_attn(
        self,
        input_batch,
        cudagraph_mode,
        block_tables,
        slot_mappings,
        attn_groups,
        kv_cache_config,
        for_capture=False,
        ubatch_idx: int = 0,
    ) -> dict[str, Any]:
        assert ubatch_idx == 0, "DBO is not supported"
        if cudagraph_mode == CUDAGraphMode.FULL:
            num_reqs = input_batch.num_reqs_after_padding
            num_tokens = input_batch.num_tokens_after_padding
        else:
            num_reqs = input_batch.num_reqs
            num_tokens = input_batch.num_tokens

        query_start_loc_cpu = torch.from_numpy(input_batch.query_start_loc_np)
        max_query_len = input_batch.num_scheduled_tokens.max().item()

        # Per-request causal mode: encoder (commit) = causal,
        # denoise = bidirectional. Pass GPU tensor so the attention
        # backend can handle mixed batches.
        actual_num_reqs = input_batch.num_reqs
        slots = input_batch.idx_mapping[:actual_num_reqs]
        # Invariant: the sampler flips is_encoder_phase to False only after a
        # request's FINAL prompt chunk, so a prompt spanning multiple chunks
        # (longer than the token budget) stays causal for every chunk.
        self._causal_buf[:actual_num_reqs].copy_(
            self.diffusion_states.is_encoder_phase[slots]
        )
        if actual_num_reqs < num_reqs:
            self._causal_buf[actual_num_reqs:num_reqs] = 0
        causal: bool | torch.Tensor = self._causal_buf[:num_reqs]

        return build_attn_metadata(
            attn_groups=attn_groups,
            num_reqs=num_reqs,
            num_tokens=num_tokens,
            query_start_loc_gpu=input_batch.query_start_loc,
            query_start_loc_cpu=query_start_loc_cpu,
            max_query_len=max_query_len,
            seq_lens=input_batch.seq_lens,
            max_seq_len=self.max_model_len,
            block_tables=block_tables,
            slot_mappings=slot_mappings,
            kv_cache_config=kv_cache_config,
            causal=causal,
        )

    num_new_sampled_tokens_per_step: int = 0

DiffusionGemmaProcessingInfo

Bases: Gemma4ProcessingInfo

Processing info for DiffusionGemma.

Overrides get_hf_config to accept DiffusionGemmaConfig (which inherits from PreTrainedConfig, not Gemma4Config). Supports image and video modalities.

Source code in vllm/model_executor/models/diffusion_gemma.py
class DiffusionGemmaProcessingInfo(Gemma4ProcessingInfo):
    """Processing info for DiffusionGemma.

    Overrides ``get_hf_config`` to accept ``DiffusionGemmaConfig``
    (which inherits from ``PreTrainedConfig``, not ``Gemma4Config``).
    Supports image and video modalities.
    """

    def get_hf_config(self):
        # DiffusionGemmaConfig doesn't inherit from Gemma4Config, so we
        # accept any PreTrainedConfig here.
        return self.ctx.get_hf_config()

    def get_supported_mm_limits(self) -> Mapping[str, int | None]:
        # DiffusionGemma supports image and video inputs.
        return {"image": None, "video": None}

    def get_mm_max_tokens_per_item(
        self, seq_len: int, mm_counts: Mapping[str, int]
    ) -> Mapping[str, int] | None:
        return super().get_mm_max_tokens_per_item(seq_len, mm_counts)

DiffusionGemmaRequestStates

Pre-allocated GPU tensors for DiffusionGemma per-request state.

Follows the indexed-slot pattern used by RequestState.

Methods:

  • allowed_tensor –

    ids as an int64 tensor on the device, built once per tuple.

  • apply_seed_canvases –

    Replace the canvas of every seeded slot among slots_gpu.

  • batch_allowed –

    The allowed ids shared by every one of slots, or None.

  • init_canvas –

    Initialize canvas with random tokens for the given slots.

  • set_pins –

    Hold positions of the slot's seed canvas through every denoise

  • set_seed_canvas –

    ids covers the slot's canvas width; positions past it are never

Source code in vllm/model_executor/models/diffusion_gemma.py
class DiffusionGemmaRequestStates:
    """Pre-allocated GPU tensors for DiffusionGemma per-request state.

    Follows the indexed-slot pattern used by ``RequestState``.
    """

    def __init__(
        self,
        max_num_reqs: int,
        canvas_length: int,
        vocab_size: int,
        max_denoising_steps: int,
        device: torch.device,
        hidden_size: int,
        stability_threshold: int,
    ):
        self.max_num_reqs = max_num_reqs
        self.canvas_length = canvas_length
        self.vocab_size = vocab_size
        self.max_denoising_steps = max_denoising_steps
        self.stability_threshold = stability_threshold
        self.device = device

        self.is_encoder_phase = torch.zeros(
            max_num_reqs, dtype=torch.bool, device=device
        )
        # Canvas tokens [max_num_reqs, canvas_length]
        self.canvas = torch.zeros(
            max_num_reqs, canvas_length, dtype=torch.int64, device=device
        )
        # Step counter (counts up from 0 to max_denoising_steps)
        self.step = torch.zeros(
            max_num_reqs,
            dtype=torch.int32,
            device=device,
        )
        # Accepted canvas history for stability check
        self.accepted_canvas_history = torch.zeros(
            max_num_reqs,
            stability_threshold,
            canvas_length,
            dtype=torch.int64,
            device=device,
        )
        self.accepted_canvas_history_len = torch.zeros(
            max_num_reqs, dtype=torch.int32, device=device
        )
        # Latest argmax(processed_logits) per slot — what we COMMIT.
        # NOT `current_canvas` (which is the post-renoise stochastic input for
        # the next denoise step). We keep this separate from `canvas` because
        # canvas gets renoised in-place during denoise, while argmax_canvas is
        # the deterministic best-guess we ultimately emit.
        self.argmax_canvas = torch.zeros(
            max_num_reqs, canvas_length, dtype=torch.int64, device=device
        )

        # Per-slot prompt length (set by add_request).
        self.prompt_len = torch.zeros(
            max_num_reqs,
            dtype=torch.int32,
            device=device,
        )

        # Per-slot confidence flag, set by the sampler each step.
        self.confident = torch.zeros(max_num_reqs, dtype=torch.bool, device=device)

        # Per-slot step cap, lowered per request from extra_args.
        self.max_steps = torch.full(
            (max_num_reqs,), max_denoising_steps, dtype=torch.int32, device=device
        )
        # Seed canvases replace the random canvas after prefill. The host-side
        # sets gate the seed and read-only work so plain generation skips it.
        self.seed_canvas = torch.zeros(
            max_num_reqs, canvas_length, dtype=torch.int64, device=device
        )
        self.has_seed = torch.zeros(max_num_reqs, dtype=torch.bool, device=device)
        self.seeded_slots: set[int] = set()
        # Pinned positions keep their seed value on every denoise step, so a
        # multi-step read denoises the canvas the request seeded.
        self.pin_mask = torch.zeros(
            max_num_reqs, canvas_length, dtype=torch.bool, device=device
        )
        # Read-only slots emit on their converging step and skip the commit
        # forward.
        self.read_only = torch.zeros(max_num_reqs, dtype=torch.bool, device=device)
        # Constrained reads: slot -> the allowed ids (the request's
        # logprob_token_ids). A step whose decode slots all share one set
        # runs the unembedding, sampler and self-conditioning over that set.
        self.constrained: dict[int, tuple[int, ...]] = {}
        self.step_allowed: torch.Tensor | None = None
        self._allowed_cache: dict[tuple[int, ...], torch.Tensor] = {}
        self.read_only_slots: set[int] = set()
        # Slots capped at one denoise step never consume a soft embed.
        self.single_step_slots: set[int] = set()
        # Per-slot canvas width, at most canvas_length. The scheduler schedules
        # this many draft tokens for the slot and the sampler pads the rest.
        self.canvas_width_np = np.full(max_num_reqs, canvas_length, dtype=np.int32)

        # Per-slot self-conditioning soft embedding (probs @ embed_weight) from
        # the previous denoise step. Storing the [.., hidden] soft embed instead
        # of the full [.., vocab] distribution shrinks this buffer by
        # vocab/hidden (~170x) and moves the matmul to denoise time; the result
        # is identical (SC consumes probs @ embed_weight anyway).
        self.self_conditioning_embeds = torch.zeros(
            max_num_reqs, canvas_length, hidden_size, dtype=torch.float32, device=device
        )

    def init_canvas(self, slot_indices: torch.Tensor) -> None:
        """Initialize canvas with random tokens for the given slots.

        `slot_indices` must already be on device to avoid a cpu->gpu sync.
        """
        n = slot_indices.shape[0]
        self.canvas[slot_indices] = torch.randint(
            0,
            self.vocab_size,
            (n, self.canvas_length),
            dtype=torch.int64,
            device=self.device,
        )

    def add_request(self, slot_idx: int) -> None:
        self.is_encoder_phase[slot_idx].fill_(True)
        self.init_canvas(async_tensor_h2d([slot_idx], device=self.device))
        self.step[slot_idx].fill_(0)
        self.accepted_canvas_history_len[slot_idx].fill_(0)
        self.self_conditioning_embeds[slot_idx] = 0
        self.max_steps[slot_idx].fill_(self.max_denoising_steps)
        self.has_seed[slot_idx].fill_(False)
        self.seeded_slots.discard(slot_idx)
        self.pin_mask[slot_idx].fill_(False)
        self.read_only[slot_idx].fill_(False)
        self.read_only_slots.discard(slot_idx)
        self.single_step_slots.discard(slot_idx)
        self.constrained.pop(slot_idx, None)
        self.canvas_width_np[slot_idx] = self.canvas_length

    def remove_request(self, slot_idx: int) -> None:
        # add_request resets the GPU flags before a slot is reused. The host
        # sets gate work on every step, so they clear here.
        self.is_encoder_phase[slot_idx].fill_(False)
        self.accepted_canvas_history_len[slot_idx].fill_(0)
        self.self_conditioning_embeds[slot_idx] = 0
        self.seeded_slots.discard(slot_idx)
        self.read_only_slots.discard(slot_idx)
        self.single_step_slots.discard(slot_idx)
        self.constrained.pop(slot_idx, None)

    def set_seed_canvas(self, slot_idx: int, ids: list[int]) -> None:
        """``ids`` covers the slot's canvas width; positions past it are never
        scheduled."""
        self.seed_canvas[slot_idx, : len(ids)] = async_tensor_h2d(
            ids, dtype=torch.int64, device=self.device
        )
        self.has_seed[slot_idx].fill_(True)
        self.seeded_slots.add(slot_idx)

    def set_pins(self, slot_idx: int, positions: list[int]) -> None:
        """Hold ``positions`` of the slot's seed canvas through every denoise
        step. Validated against the request's canvas width upstream."""
        self.pin_mask[slot_idx].fill_(False)
        self.pin_mask[
            slot_idx, async_tensor_h2d(positions, dtype=torch.int64, device=self.device)
        ] = True

    def allowed_tensor(self, ids: tuple[int, ...]) -> torch.Tensor:
        """``ids`` as an int64 tensor on the device, built once per tuple."""
        t = self._allowed_cache.get(ids)
        if t is None:
            t = torch.tensor(ids, dtype=torch.int64, device=self.device)
            if len(self._allowed_cache) > 64:
                self._allowed_cache.clear()
            self._allowed_cache[ids] = t
        return t

    def batch_allowed(self, slots: list[int]) -> torch.Tensor | None:
        """The allowed ids shared by every one of ``slots``, or None.

        None means the step runs over the full vocabulary. The sampler then
        masks each constrained slot's logit rows to its own set, so a
        constrained request reads the same way whoever shares its batch.
        """
        if not slots or not self.constrained:
            return None
        first = self.constrained.get(slots[0])
        if first is None or any(self.constrained.get(s) != first for s in slots[1:]):
            return None
        return self.allowed_tensor(first)

    def set_read_only(self, slot_idx: int) -> None:
        self.read_only[slot_idx].fill_(True)
        self.read_only_slots.add(slot_idx)

    def apply_seed_canvases(
        self, slots_np: np.ndarray, slots_gpu: torch.Tensor
    ) -> None:
        """Replace the canvas of every seeded slot among ``slots_gpu``."""
        if self.seeded_slots.isdisjoint(slots_np.tolist()):
            return
        self.canvas[slots_gpu] = torch.where(
            self.has_seed[slots_gpu, None],
            self.seed_canvas[slots_gpu],
            self.canvas[slots_gpu],
        )

allowed_tensor(ids)

ids as an int64 tensor on the device, built once per tuple.

Source code in vllm/model_executor/models/diffusion_gemma.py
def allowed_tensor(self, ids: tuple[int, ...]) -> torch.Tensor:
    """``ids`` as an int64 tensor on the device, built once per tuple."""
    t = self._allowed_cache.get(ids)
    if t is None:
        t = torch.tensor(ids, dtype=torch.int64, device=self.device)
        if len(self._allowed_cache) > 64:
            self._allowed_cache.clear()
        self._allowed_cache[ids] = t
    return t

apply_seed_canvases(slots_np, slots_gpu)

Replace the canvas of every seeded slot among slots_gpu.

Source code in vllm/model_executor/models/diffusion_gemma.py
def apply_seed_canvases(
    self, slots_np: np.ndarray, slots_gpu: torch.Tensor
) -> None:
    """Replace the canvas of every seeded slot among ``slots_gpu``."""
    if self.seeded_slots.isdisjoint(slots_np.tolist()):
        return
    self.canvas[slots_gpu] = torch.where(
        self.has_seed[slots_gpu, None],
        self.seed_canvas[slots_gpu],
        self.canvas[slots_gpu],
    )

batch_allowed(slots)

The allowed ids shared by every one of slots, or None.

None means the step runs over the full vocabulary. The sampler then masks each constrained slot's logit rows to its own set, so a constrained request reads the same way whoever shares its batch.

Source code in vllm/model_executor/models/diffusion_gemma.py
def batch_allowed(self, slots: list[int]) -> torch.Tensor | None:
    """The allowed ids shared by every one of ``slots``, or None.

    None means the step runs over the full vocabulary. The sampler then
    masks each constrained slot's logit rows to its own set, so a
    constrained request reads the same way whoever shares its batch.
    """
    if not slots or not self.constrained:
        return None
    first = self.constrained.get(slots[0])
    if first is None or any(self.constrained.get(s) != first for s in slots[1:]):
        return None
    return self.allowed_tensor(first)

init_canvas(slot_indices)

Initialize canvas with random tokens for the given slots.

slot_indices must already be on device to avoid a cpu->gpu sync.

Source code in vllm/model_executor/models/diffusion_gemma.py
def init_canvas(self, slot_indices: torch.Tensor) -> None:
    """Initialize canvas with random tokens for the given slots.

    `slot_indices` must already be on device to avoid a cpu->gpu sync.
    """
    n = slot_indices.shape[0]
    self.canvas[slot_indices] = torch.randint(
        0,
        self.vocab_size,
        (n, self.canvas_length),
        dtype=torch.int64,
        device=self.device,
    )

set_pins(slot_idx, positions)

Hold positions of the slot's seed canvas through every denoise step. Validated against the request's canvas width upstream.

Source code in vllm/model_executor/models/diffusion_gemma.py
def set_pins(self, slot_idx: int, positions: list[int]) -> None:
    """Hold ``positions`` of the slot's seed canvas through every denoise
    step. Validated against the request's canvas width upstream."""
    self.pin_mask[slot_idx].fill_(False)
    self.pin_mask[
        slot_idx, async_tensor_h2d(positions, dtype=torch.int64, device=self.device)
    ] = True

set_seed_canvas(slot_idx, ids)

ids covers the slot's canvas width; positions past it are never scheduled.

Source code in vllm/model_executor/models/diffusion_gemma.py
def set_seed_canvas(self, slot_idx: int, ids: list[int]) -> None:
    """``ids`` covers the slot's canvas width; positions past it are never
    scheduled."""
    self.seed_canvas[slot_idx, : len(ids)] = async_tensor_h2d(
        ids, dtype=torch.int64, device=self.device
    )
    self.has_seed[slot_idx].fill_(True)
    self.seeded_slots.add(slot_idx)

DiffusionGemmaSelfConditioning

Bases: Module

Gated MLP that processes soft embeddings from the previous denoising step.

Structurally identical to Gemma4MLP but with self_conditioning_size and post_norm without learned scale.

Source code in vllm/model_executor/models/diffusion_gemma.py
class DiffusionGemmaSelfConditioning(nn.Module):
    """Gated MLP that processes soft embeddings from the previous denoising step.

    Structurally identical to Gemma4MLP but with self_conditioning_size
    and post_norm without learned scale.
    """

    def __init__(
        self, hidden_size: int, self_conditioning_size: int, eps: float = 1e-6
    ):
        super().__init__()
        self.pre_norm = RMSNorm(hidden_size, eps=eps)
        self.post_norm = RMSNorm(hidden_size, eps=eps, has_weight=False)
        self.gate_proj = nn.Linear(hidden_size, self_conditioning_size, bias=False)
        self.up_proj = nn.Linear(hidden_size, self_conditioning_size, bias=False)
        self.down_proj = nn.Linear(self_conditioning_size, hidden_size, bias=False)

    def forward(
        self,
        inputs_embeds: torch.Tensor,
        soft_embeds: torch.Tensor,
    ) -> torch.Tensor:
        x = self.pre_norm(soft_embeds)
        sc_signal = self.down_proj(
            F.gelu(self.gate_proj(x), approximate="tanh") * self.up_proj(x)
        )
        return self.post_norm(inputs_embeds + sc_signal)

DiffusionSampler

Batched accept/renoise sampler for DiffusionGemma.

Follows the same structure as vllm.v1.worker.gpu.sample.sampler.Sampler: decomposed into named methods, all GPU state in pre-allocated buffers, no GPU→CPU syncs on the hot path.

Source code in vllm/model_executor/models/diffusion_gemma.py
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243
1244
1245
1246
1247
1248
1249
1250
1251
1252
1253
1254
1255
1256
1257
1258
1259
1260
1261
1262
1263
1264
1265
1266
1267
1268
1269
1270
1271
1272
1273
1274
1275
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
1287
1288
1289
1290
1291
1292
1293
1294
1295
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
1319
1320
1321
1322
1323
1324
1325
1326
1327
1328
1329
1330
1331
1332
1333
1334
1335
1336
1337
1338
1339
1340
1341
1342
1343
1344
1345
1346
1347
1348
1349
1350
1351
1352
1353
1354
1355
1356
1357
1358
1359
1360
1361
1362
1363
1364
1365
1366
1367
1368
1369
1370
1371
1372
1373
1374
1375
1376
1377
1378
1379
1380
1381
1382
1383
1384
1385
1386
1387
1388
1389
1390
1391
1392
1393
1394
1395
1396
1397
1398
1399
1400
1401
1402
1403
1404
1405
1406
1407
1408
1409
1410
1411
1412
1413
1414
1415
1416
1417
1418
1419
1420
1421
1422
1423
1424
1425
1426
1427
1428
1429
1430
1431
1432
1433
1434
1435
1436
1437
1438
1439
1440
1441
1442
1443
1444
1445
1446
1447
1448
1449
1450
1451
1452
1453
1454
1455
1456
1457
1458
1459
1460
1461
1462
1463
1464
1465
1466
1467
1468
1469
1470
1471
1472
1473
1474
1475
1476
1477
1478
1479
1480
1481
1482
1483
1484
1485
1486
1487
1488
1489
1490
1491
1492
1493
1494
1495
1496
1497
1498
1499
1500
1501
1502
1503
1504
1505
1506
1507
1508
1509
1510
1511
1512
1513
1514
1515
1516
1517
1518
1519
1520
1521
1522
1523
1524
1525
1526
1527
1528
1529
1530
1531
1532
1533
1534
1535
1536
1537
1538
1539
1540
1541
1542
1543
1544
1545
1546
1547
1548
1549
1550
1551
1552
1553
1554
1555
1556
1557
1558
1559
1560
1561
1562
1563
1564
1565
1566
1567
1568
1569
1570
1571
1572
1573
1574
1575
1576
1577
1578
1579
1580
1581
1582
1583
1584
1585
1586
1587
1588
1589
1590
1591
1592
1593
1594
1595
1596
1597
1598
1599
1600
1601
1602
1603
1604
1605
1606
1607
1608
1609
1610
1611
1612
1613
1614
1615
1616
1617
1618
1619
1620
1621
1622
1623
1624
1625
1626
1627
1628
1629
1630
1631
1632
1633
1634
1635
1636
1637
1638
1639
1640
1641
1642
1643
1644
1645
1646
1647
1648
1649
1650
1651
1652
1653
1654
1655
1656
1657
1658
1659
1660
1661
1662
1663
1664
1665
1666
1667
1668
1669
1670
1671
1672
1673
1674
1675
1676
1677
1678
1679
1680
1681
1682
1683
1684
1685
1686
1687
1688
1689
1690
1691
1692
1693
1694
1695
1696
1697
1698
1699
1700
1701
1702
1703
1704
1705
1706
1707
1708
1709
class DiffusionSampler:
    """Batched accept/renoise sampler for DiffusionGemma.

    Follows the same structure as ``vllm.v1.worker.gpu.sample.sampler.Sampler``:
    decomposed into named methods, all GPU state in pre-allocated buffers,
    no GPU→CPU syncs on the hot path.
    """

    def __init__(
        self,
        sampler: Any,
        diffusion_config: Any,
        vocab_size: int,
        diffusion_states: DiffusionGemmaRequestStates,
        *,
        confidence_threshold: float,
        t_min: float,
        t_max: float,
        entropy_bound: float,
        embed_weight: torch.Tensor,
        normalizer: torch.Tensor,
        sc_vocab_start: int = 0,
        sc_vocab_end: int | None = None,
        tp_size: int = 1,
        tp_group_name: str = "",
    ):
        self.sampling_states = sampler.sampling_states
        self.logprob_token_ids_state = sampler.logprob_token_ids_state
        self.req_states = sampler.req_states
        self.logits_mode = sampler.logprobs_mode in ("raw_logits", "processed_logits")
        # Self-conditioning soft embed = probs @ embed_weight * normalizer,
        # computed in the sampler (see _compiled_sample_step). ``embed_weight``
        # is the vocab-parallel shard; [sc_vocab_start, sc_vocab_end) is this
        # rank's slice of the full vocab and tp_* drive the cross-rank
        # all-reduce.
        self.embed_weight = embed_weight
        self.normalizer = normalizer
        self.sc_vocab_start = sc_vocab_start
        self.sc_vocab_end = sc_vocab_end if sc_vocab_end is not None else vocab_size
        self.tp_size = tp_size
        self.tp_group_name = tp_group_name
        self.canvas_length = (
            diffusion_config.canvas_length if diffusion_config is not None else 32
        )
        self.t_min = t_min
        self.t_max = t_max
        self.confidence_threshold = confidence_threshold
        self.vocab_size = vocab_size
        self.diffusion_states = diffusion_states
        self.entropy_bound = entropy_bound

        max_num_reqs = diffusion_states.max_num_reqs
        device = diffusion_states.device
        self._sampled = torch.zeros(
            max_num_reqs,
            self.canvas_length,
            dtype=torch.int32,
            device=device,
        )
        self._num_sampled = torch.zeros(
            max_num_reqs,
            dtype=torch.int32,
            device=device,
        )
        self._decode_slots = UvaBackedTensor(max_num_reqs, dtype=torch.int64)
        self._decode_idx = UvaBackedTensor(max_num_reqs, dtype=torch.int64)
        self._query_lens = UvaBackedTensor(max_num_reqs, dtype=torch.int32)
        self._num_logits = UvaBackedTensor(max_num_reqs, dtype=torch.int32)

        # Per-slot stash for logprobs computed on the converging denoise step.
        # Populated after the post-sample kernel detects convergence; consumed
        # on the subsequent commit step when num_sampled=CANVAS_LEN.
        self._pending_logprobs: dict[int, LogprobsTensors] = {}

    def add_request(self, req_idx: int, sampling_params: Any) -> None:
        if use_penalty(sampling_params):
            logger.warning_once(
                "DiffusionGemma does not support repetition/frequency/presence "
                "penalties; ignoring them for this request."
            )
        # Purge any stale logprobs stashed under this slot by a prior request
        # that was aborted between its converging denoise and commit steps.
        self._pending_logprobs.pop(req_idx, None)
        self.sampling_states.add_request(req_idx, sampling_params)
        self.logprob_token_ids_state.add_request(req_idx, sampling_params)
        extra = getattr(sampling_params, "extra_args", None) or {}
        states = self.diffusion_states
        cap = extra.get("diffusion_max_steps")
        if cap is not None:
            cap = max(1, min(int(cap), states.max_denoising_steps))
            states.max_steps[req_idx].fill_(cap)
            if cap == 1:
                states.single_step_slots.add(req_idx)
        width = extra.get("diffusion_canvas_length")
        if width:
            states.canvas_width_np[req_idx] = max(
                1, min(int(width), self.canvas_length)
            )
        width = int(states.canvas_width_np[req_idx])
        seed = extra.get("diffusion_seed_canvas")
        if seed is not None:
            if len(seed) != width:
                raise ValueError(
                    f"diffusion_seed_canvas must hold exactly {width} ids, "
                    f"got {len(seed)}"
                )
            states.set_seed_canvas(req_idx, seed)
        pins = extra.get("diffusion_pinned")
        if pins and seed is not None:
            states.set_pins(req_idx, [int(p) for p in pins])
        if extra.get("diffusion_read_only"):
            states.set_read_only(req_idx)
        if extra.get("diffusion_constrained"):
            ids = list(getattr(sampling_params, "logprob_token_ids", None) or [])
            if not ids:
                raise ValueError(
                    "diffusion_constrained needs logprob_token_ids: they are "
                    "the allowed set."
                )
            if self.tp_size > 1:
                raise ValueError("diffusion_constrained needs tensor parallel 1.")
            states.constrained[req_idx] = tuple(int(t) for t in ids)

    def apply_staged_writes(self) -> None:
        self.sampling_states.apply_staged_writes()
        self.logprob_token_ids_state.apply_staged_writes()

    @property
    def penalties_state(self):
        # Diffusion applies no penalties. The runner reads
        # penalties_state.output_bin_counts, so expose a stub holding None;
        # post_update treats None bin counts as "no penalty bookkeeping".
        return _NO_PENALTIES_STATE

    # ------------------------------------------------------------------
    # Prefill
    # ------------------------------------------------------------------

    def _finish_prefills(
        self, input_batch: Any, prefill_indices_np: np.ndarray
    ) -> None:
        """Transition requests whose prompt completes this step to denoising.

        Initializes their canvas, seeds draft tokens, and flips
        is_encoder_phase to False. Mid-chunk requests (prompt longer than the
        token budget) are left untouched so is_encoder_phase stays True and
        prepare_attn keeps causal attention for their remaining chunks.
        """
        states = self.diffusion_states
        done_prefill_np = (
            input_batch.num_computed_prefill_tokens_np[prefill_indices_np]
            + input_batch.num_scheduled_tokens[prefill_indices_np]
            >= input_batch.prefill_len_np[prefill_indices_np]
        )
        ps = input_batch.idx_mapping_np[prefill_indices_np[done_prefill_np]]
        if len(ps) == 0:
            return
        # Move the slot indices across once, up front: indexing a device
        # tensor with a numpy array copies them over synchronously each time.
        ps_gpu = async_tensor_h2d(
            ps.astype(np.int64), device=states.is_encoder_phase.device
        )
        states.init_canvas(ps_gpu)
        states.apply_seed_canvases(ps, ps_gpu)
        self.req_states.draft_tokens[ps_gpu, : self.canvas_length] = states.canvas[
            ps_gpu
        ]
        states.is_encoder_phase.index_fill_(0, ps_gpu, False)

    def _handle_prefill(
        self,
        input_batch: Any,
        device: torch.device,
    ) -> SamplerOutput:
        num_reqs = input_batch.num_reqs
        self._finish_prefills(input_batch, np.arange(num_reqs))
        sampled = self._sampled[:num_reqs, :1]
        sampled.zero_()
        num_sampled = self._num_sampled[:num_reqs]
        num_sampled.zero_()
        return SamplerOutput(
            sampled_token_ids=sampled,
            logprobs_tensors=None,
            num_nans=None,
            num_sampled=num_sampled,
            num_rejected=num_sampled,
        )

    # ------------------------------------------------------------------
    # Decode helpers
    # ------------------------------------------------------------------

    def _build_output(
        self,
        input_batch: Any,
        sampled: torch.Tensor,
        num_sampled: torch.Tensor,
        per_req_nlogits_np: np.ndarray,
        device: torch.device,
        logprobs_tensors: LogprobsTensors | None = None,
    ) -> SamplerOutput:
        """Compute num_rejected and build SamplerOutput."""
        num_reqs = input_batch.num_reqs

        self._query_lens.np[:num_reqs] = np.diff(
            input_batch.query_start_loc_np[: num_reqs + 1]
        )
        self._num_logits.np[:num_reqs] = per_req_nlogits_np
        self._query_lens.copy_to_uva()
        self._num_logits.copy_to_uva()

        num_rejected = _compute_num_rejected(
            self._num_logits.gpu[:num_reqs],
            num_sampled,
            input_batch.query_start_loc[: num_reqs + 1],
        )

        return SamplerOutput(
            sampled_token_ids=sampled,
            logprobs_tensors=logprobs_tensors,
            num_nans=None,
            num_sampled=num_sampled,
            num_rejected=num_rejected,
        )

    # ------------------------------------------------------------------
    # Main entry point
    # ------------------------------------------------------------------

    def __call__(
        self,
        logits: torch.Tensor,
        input_batch: Any,
        draft_logits: torch.Tensor | None = None,
    ) -> SamplerOutput:
        num_reqs = input_batch.num_reqs
        device = logits.device

        if input_batch.num_draft_tokens == 0:
            return self._handle_prefill(input_batch, device)

        # --- CPU/NumPy setup (outside compile): split decode vs prefill, init
        # canvas for any new prefills, and stage decode slot indices to GPU. ---
        states = self.diffusion_states
        slots_np = input_batch.idx_mapping_np[:num_reqs]
        per_req_nlogits_np = np.diff(input_batch.cu_num_logits_np[: num_reqs + 1])

        decode_indices_np = np.where(per_req_nlogits_np > 0)[0]
        prefill_indices_np = np.where(per_req_nlogits_np == 0)[0]
        decode_slots_np = slots_np[decode_indices_np]

        if len(prefill_indices_np) > 0:
            self._finish_prefills(input_batch, prefill_indices_np)

        num_decode = len(decode_indices_np)
        self._decode_slots.np[:num_decode] = decode_slots_np
        self._decode_idx.np[:num_decode] = decode_indices_np
        self._decode_slots.copy_to_uva()
        self._decode_idx.copy_to_uva()
        decode_slots = self._decode_slots.gpu[:num_decode]
        decode_idx = self._decode_idx.gpu[:num_decode]

        # Real canvas length per decode request. Equals CL except when a canvas
        # was truncated near max_model_len, in which case the scheduler gave us
        # fewer than CL logits for that request.
        valid_canvas_len_np = per_req_nlogits_np[per_req_nlogits_np > 0]
        valid_canvas_len = async_tensor_h2d(
            valid_canvas_len_np.astype(np.int64), device=device
        )

        # Per-request top_k/top_p, mirroring the AR sampler. Masked tokens
        # become -inf and survive the temperature scaling in the compiled
        # step, so Gumbel sampling, probs, and entropy all see the filtered
        # distribution. The committed argmax (always the top-1 token) is
        # unaffected; only the canvas exploration is constrained. Applied
        # before canvas padding so phantom positions stay uniform.
        if num_decode > 0:
            top_k, top_p = self.sampling_states.get_top_k_top_p(
                decode_slots.repeat_interleave(
                    valid_canvas_len, output_size=int(valid_canvas_len_np.sum())
                ),
                decode_slots_np,
            )
            if top_k is not None or top_p is not None:
                logits = apply_top_k_top_p(logits.float(), top_k, top_p)

        # Where each decode request's rows start in the flat logits. Tiles
        # below gather and pad a request to the tile's width.
        row_starts_np = np.concatenate(
            ([0], np.cumsum(valid_canvas_len_np)[:-1])
        ).astype(np.int64)
        row_starts = async_tensor_h2d(row_starts_np, device=device)

        # Clear once: the tiled loop below only scatters its own decode slots,
        # so it must not re-clear earlier tiles' writes.
        sampled = self._sampled[:num_reqs]
        num_sampled = self._num_sampled[:num_reqs]
        sampled.zero_()
        num_sampled.zero_()

        all_slots = input_batch.idx_mapping[:num_reqs]

        # Snapshot which slots are committing BEFORE the compiled step runs,
        # since it mutates is_encoder_phase (commit→False, converge→True).
        is_committing = states.is_encoder_phase[decode_slots].clone()

        # Constrained step: logits are [rows, K] over the shared allowed set,
        # so the self-conditioning matmul only needs those K embedding rows.
        allowed = states.step_allowed
        embed_weight = self.embed_weight
        vocab_size = self.vocab_size
        sc_vocab_start, sc_vocab_end = self.sc_vocab_start, self.sc_vocab_end
        if allowed is not None:
            assert logits.shape[-1] == allowed.numel()
            embed_weight = self.embed_weight[allowed]
            sc_vocab_start, sc_vocab_end = 0, allowed.numel()
        elif states.constrained and num_decode > 0:
            # The decode slots do not share one set, so this step runs over
            # the full vocabulary. Mask each constrained request's rows to
            # its own set so it reads the same as on the shared path. The
            # helper masks a copy, since the runner owns `logits`: one
            # full-vocab copy per step, as for top_k/top_p above.
            per_row = [
                None if ids is None else states.allowed_tensor(ids)
                for ids in (states.constrained.get(s) for s in decode_slots_np.tolist())
            ]
            logits = _mask_rows_to_allowed(
                logits, row_starts_np, valid_canvas_len_np, per_row
            )

        slots_np = input_batch.idx_mapping_np[:num_reqs]
        max_num_logprobs = self.sampling_states.max_num_logprobs(slots_np)
        # Requests may ask for specific token ids' logprobs instead of, or as
        # well as, a top-k.
        max_token_ids = self.logprob_token_ids_state.max_num_token_ids(slots_np)
        want_logprobs = max_num_logprobs >= 0 or max_token_ids > 0
        num_logprobs = max(max_num_logprobs, 0)

        # Sample per tile. Decode requests are grouped by canvas width and each
        # tile runs the compiled step at that width over [:, :W] views of the
        # state, so a narrow read pays for its own rows rather than the served
        # canvas. Widths ascend, so the last tile's canvas-to-draft copy (over
        # all slots) is the widest. The fp32 pipeline keeps several live
        # [tile * W, vocab] copies, so a tile is also bounded by free memory.
        widths_np = states.canvas_width_np[decode_slots_np]
        order = np.argsort(widths_np, kind="stable")
        free = torch.accelerator.get_memory_info()[0] if num_decode > 0 else 0
        run_start = 0
        while run_start < num_decode:
            W = int(widths_np[order[run_start]])
            run_end = run_start
            while run_end < num_decode and widths_np[order[run_end]] == W:
                run_end += 1
            # Transient [tile * W, vocab] tensors: the softmax in the embedding
            # dtype for self-conditioning, plus the fp32 scaled logits when
            # logprobs are wanted (pad for allocator overhead).
            budget = max(1, int(free * 0.5) // max(W * self.vocab_size * 4 * 3, 1))
            for t0 in range(run_start, run_end, budget):
                sel_np = order[t0 : min(t0 + budget, run_end)]
                n = len(sel_np)
                sel = async_tensor_h2d(sel_np.astype(np.int64), device=device)
                tile_slots = decode_slots[sel]
                tile_valid = valid_canvas_len[sel]
                tile_valid_np = valid_canvas_len_np[sel_np]
                contiguous = bool((sel_np == np.arange(sel_np[0], sel_np[0] + n)).all())
                if contiguous and tile_valid_np.min() == W:
                    r0 = int(row_starts_np[sel_np[0]])
                    tile_logits = logits[r0 : r0 + n * W]
                else:
                    # Pad each request to W. Phantom positions are zeroed:
                    # uniform logits, high entropy, argmax 0, never committed.
                    # masked_fill, not multiply, so -inf from top_k/top_p
                    # filtering above cannot turn a phantom row into NaN.
                    ar = torch.arange(W, device=device)
                    src = (row_starts[sel].unsqueeze(1) + ar.unsqueeze(0)).clamp_max(
                        logits.shape[0] - 1
                    )
                    valid = ar.unsqueeze(0) < tile_valid.unsqueeze(1)
                    tile_logits = logits[src.reshape(-1)].masked_fill_(
                        ~valid.reshape(-1, 1), 0
                    )
                compute_sc = (
                    not states.single_step_slots
                    or not states.single_step_slots.issuperset(
                        decode_slots_np[sel_np].tolist()
                    )
                )

                temp = _denoise_temperature(
                    states.step,
                    tile_slots,
                    float(states.max_denoising_steps),
                    self.t_min,
                    self.t_max,
                )
                probs_dtype = self.embed_weight.dtype if compute_sc else None
                if tile_logits.is_cuda:
                    seed = int(torch.randint(0, 2**31 - 1, (1,)).item())
                    stats = sample_row_stats(tile_logits, temp, W, seed, probs_dtype)
                else:
                    stats = sample_row_stats_reference(
                        tile_logits, temp, W, probs_dtype
                    )
                argmax_rows, sample_rows, entropy_rows, probs_rows = stats
                if allowed is not None:
                    # K-space picks back to token ids. The softmax stays K wide
                    # for the K-row self-conditioning matmul.
                    argmax_rows = allowed[argmax_rows]
                    sample_rows = allowed[sample_rows]
                probs = probs_rows.view(n, W, -1) if probs_rows is not None else None
                # Only the logprob stash reads the scaled logits.
                scaled = None
                if want_logprobs:
                    scaled = tile_logits.float().view(n, W, -1) / temp[
                        :, None, None
                    ].clamp(min=1e-10)

                _compiled_sample_step(
                    sample_rows.view(n, W),
                    argmax_rows.view(n, W),
                    entropy_rows.view(n, W),
                    probs,
                    tile_slots,
                    decode_idx[sel],
                    all_slots,
                    tile_valid,
                    # State, viewed at this tile's width
                    states.canvas[:, :W],
                    states.argmax_canvas[:, :W],
                    states.step,
                    states.is_encoder_phase,
                    states.confident,
                    states.self_conditioning_embeds[:, :W],
                    embed_weight,
                    self.normalizer,
                    states.accepted_canvas_history[:, :, :W],
                    states.accepted_canvas_history_len,
                    states.max_steps,
                    states.pin_mask[:, :W],
                    states.seed_canvas[:, :W],
                    states.read_only,
                    # Output
                    sampled[:, :W],
                    num_sampled,
                    self.req_states.draft_tokens,
                    # Config
                    confidence_threshold=self.confidence_threshold,
                    vocab_size=vocab_size,
                    CL=W,
                    ST=states.stability_threshold,
                    entropy_bound=self.entropy_bound,
                    sc_vocab_start=sc_vocab_start,
                    sc_vocab_end=sc_vocab_end,
                    tp_size=self.tp_size,
                    tp_group_name=self.tp_group_name,
                    compute_sc=compute_sc,
                )

                # Stash newly converged logprobs, including reads that emit now.
                if want_logprobs:
                    converged_mask = states.is_encoder_phase[tile_slots] | (
                        num_sampled[decode_idx[sel]] > 0
                    )
                    just_converged = converged_mask & ~is_committing[sel]
                    if just_converged.any():
                        assert scaled is not None
                        flat_logits = scaled.reshape(-1, scaled.shape[-1])
                        argmax_tokens = scaled.argmax(dim=-1)
                        raw_flat: torch.Tensor | None = None
                        for local_idx in just_converged.nonzero(as_tuple=True)[0]:
                            li = local_idx.item()
                            slot = tile_slots[local_idx].item()
                            # Stash only the real canvas positions; padded tail
                            # positions are never emitted.
                            k_i = int(tile_valid_np[li])
                            pos = li * W
                            src = flat_logits
                            if slot in states.read_only_slots:
                                # Read-only slots report logprobs at temperature 1.
                                # The schedule-tempered logits share the argmax.
                                if raw_flat is None:
                                    raw_flat = tile_logits.float()
                                src = raw_flat
                            if allowed is not None:
                                # Column j of the K-space logits is allowed[j]:
                                # report every allowed id, renormalized over
                                # the set, with the argmax in column 0.
                                rows = src[pos : pos + k_i]
                                lp = rows.log_softmax(dim=-1)
                                am = argmax_tokens[local_idx][:k_i]
                                sel_lp = lp.gather(1, am.unsqueeze(1))
                                # Same convention as _ranks_kernel: 1-based,
                                # ties count, so the rank is the number of
                                # columns at or above the selected one. Out
                                # of set ids have zero probability, so the
                                # rank within the set is the real rank.
                                ranks = (lp >= sel_lp).sum(dim=1)
                                self._pending_logprobs[slot] = LogprobsTensors(
                                    logprob_token_ids=torch.cat(
                                        (
                                            allowed[am].unsqueeze(1),
                                            allowed.unsqueeze(0).expand(k_i, -1),
                                        ),
                                        dim=1,
                                    ),
                                    logprobs=torch.cat((sel_lp, lp), dim=1),
                                    selected_token_ranks=ranks.to(torch.int64),
                                )
                                continue
                            per_req_ids = max_token_ids > 0
                            self._pending_logprobs[slot] = compute_topk_scores(
                                src[pos : pos + k_i],
                                num_logprobs,
                                argmax_tokens[local_idx][:k_i],
                                logprob_token_ids_state=(
                                    self.logprob_token_ids_state
                                    if per_req_ids
                                    else None
                                ),
                                # every row of this stash belongs to one slot
                                expanded_idx_mapping=(
                                    torch.full(
                                        (k_i,), slot, dtype=torch.int32, device=device
                                    )
                                    if per_req_ids
                                    else None
                                ),
                                max_per_req_token_ids=max_token_ids,
                                logits_mode=self.logits_mode,
                            )
            run_start = run_end

        # Only emitting requests consume their stashed logprobs.
        logprobs_tensors = None
        if want_logprobs and self._pending_logprobs:
            emitting_slots = set(slots_np[num_sampled.cpu().numpy() > 0].tolist())
            parts: list[LogprobsTensors] = []
            cu_gen: list[int] = []
            flat_offset = 0
            for i in range(num_reqs):
                cu_gen.append(flat_offset)
                slot = int(slots_np[i])
                if slot in emitting_slots and slot in self._pending_logprobs:
                    lp = self._pending_logprobs.pop(slot)
                    parts.append(lp)
                    flat_offset += lp.logprobs.shape[0]
            if parts:
                logprobs_tensors = _concat_logprob_stashes(parts, cu_gen)

        return self._build_output(
            input_batch,
            sampled,
            num_sampled,
            per_req_nlogits_np,
            device,
            logprobs_tensors=logprobs_tensors,
        )

_build_output(input_batch, sampled, num_sampled, per_req_nlogits_np, device, logprobs_tensors=None)

Compute num_rejected and build SamplerOutput.

Source code in vllm/model_executor/models/diffusion_gemma.py
def _build_output(
    self,
    input_batch: Any,
    sampled: torch.Tensor,
    num_sampled: torch.Tensor,
    per_req_nlogits_np: np.ndarray,
    device: torch.device,
    logprobs_tensors: LogprobsTensors | None = None,
) -> SamplerOutput:
    """Compute num_rejected and build SamplerOutput."""
    num_reqs = input_batch.num_reqs

    self._query_lens.np[:num_reqs] = np.diff(
        input_batch.query_start_loc_np[: num_reqs + 1]
    )
    self._num_logits.np[:num_reqs] = per_req_nlogits_np
    self._query_lens.copy_to_uva()
    self._num_logits.copy_to_uva()

    num_rejected = _compute_num_rejected(
        self._num_logits.gpu[:num_reqs],
        num_sampled,
        input_batch.query_start_loc[: num_reqs + 1],
    )

    return SamplerOutput(
        sampled_token_ids=sampled,
        logprobs_tensors=logprobs_tensors,
        num_nans=None,
        num_sampled=num_sampled,
        num_rejected=num_rejected,
    )

_finish_prefills(input_batch, prefill_indices_np)

Transition requests whose prompt completes this step to denoising.

Initializes their canvas, seeds draft tokens, and flips is_encoder_phase to False. Mid-chunk requests (prompt longer than the token budget) are left untouched so is_encoder_phase stays True and prepare_attn keeps causal attention for their remaining chunks.

Source code in vllm/model_executor/models/diffusion_gemma.py
def _finish_prefills(
    self, input_batch: Any, prefill_indices_np: np.ndarray
) -> None:
    """Transition requests whose prompt completes this step to denoising.

    Initializes their canvas, seeds draft tokens, and flips
    is_encoder_phase to False. Mid-chunk requests (prompt longer than the
    token budget) are left untouched so is_encoder_phase stays True and
    prepare_attn keeps causal attention for their remaining chunks.
    """
    states = self.diffusion_states
    done_prefill_np = (
        input_batch.num_computed_prefill_tokens_np[prefill_indices_np]
        + input_batch.num_scheduled_tokens[prefill_indices_np]
        >= input_batch.prefill_len_np[prefill_indices_np]
    )
    ps = input_batch.idx_mapping_np[prefill_indices_np[done_prefill_np]]
    if len(ps) == 0:
        return
    # Move the slot indices across once, up front: indexing a device
    # tensor with a numpy array copies them over synchronously each time.
    ps_gpu = async_tensor_h2d(
        ps.astype(np.int64), device=states.is_encoder_phase.device
    )
    states.init_canvas(ps_gpu)
    states.apply_seed_canvases(ps, ps_gpu)
    self.req_states.draft_tokens[ps_gpu, : self.canvas_length] = states.canvas[
        ps_gpu
    ]
    states.is_encoder_phase.index_fill_(0, ps_gpu, False)

_compiled_sample_step(new_tokens, argmax_tokens, token_entropy, probs, decode_slots, decode_idx, all_slots, valid_canvas_len, canvas, argmax_canvas, step_tensor, is_encoder_phase, confident_tensor, sc_embeds, embed_weight, normalizer, history, history_len_tensor, max_steps_tensor, pin_mask, seed_canvas, read_only, sampled, num_sampled, draft_tokens, confidence_threshold, vocab_size, CL, ST, entropy_bound, sc_vocab_start, sc_vocab_end, tp_size, tp_group_name, compute_sc=True)

Compiled decode step: confidence → accept/renoise → convergence, as vectorized PyTorch ops over [num_decode, CL] tensors. The per-position statistics (argmax, Gumbel-max sample, entropy, softmax) come from one pass over the logits in sample_row_stats.

Source code in vllm/model_executor/models/diffusion_gemma.py
def _compiled_sample_step(
    # Per-position statistics of the temperature-scaled logits, from
    # sample_row_stats: [num_decode, CL] each, and the softmax
    # [num_decode, CL, vocab] in the embedding dtype when compute_sc.
    new_tokens: torch.Tensor,
    argmax_tokens: torch.Tensor,
    token_entropy: torch.Tensor,
    probs: torch.Tensor | None,
    # Request mapping
    decode_slots: torch.Tensor,  # [num_decode] int64 → slot indices
    decode_idx: torch.Tensor,  # [num_decode] int64 → position in num_reqs
    all_slots: torch.Tensor,  # [num_reqs] int64 → all slot indices
    valid_canvas_len: torch.Tensor,  # [num_decode] int64 → real canvas length (<=CL)
    # State tensors (modified in-place)
    canvas: torch.Tensor,  # [max_num_reqs, CL]
    argmax_canvas: torch.Tensor,  # [max_num_reqs, CL]
    step_tensor: torch.Tensor,  # [max_num_reqs]
    is_encoder_phase: torch.Tensor,  # [max_num_reqs]
    confident_tensor: torch.Tensor,  # [max_num_reqs]
    sc_embeds: torch.Tensor,  # [max_num_reqs, CL, hidden]
    embed_weight: torch.Tensor,  # [vocab, hidden]
    normalizer: torch.Tensor,
    history: torch.Tensor,  # [max_num_reqs, ST, CL]
    history_len_tensor: torch.Tensor,  # [max_num_reqs]
    max_steps_tensor: torch.Tensor,  # [max_num_reqs] int32, per-slot step cap
    pin_mask: torch.Tensor,  # [max_num_reqs, CL] bool, positions held at the seed
    seed_canvas: torch.Tensor,  # [max_num_reqs, CL]
    read_only: torch.Tensor,  # [max_num_reqs] bool, emit without a commit forward
    # Output tensors (modified in-place)
    sampled: torch.Tensor,  # [num_reqs, CL]
    num_sampled: torch.Tensor,  # [num_reqs]
    draft_tokens: torch.Tensor,  # [max_num_reqs, >=CL]
    # Scalar config
    confidence_threshold: float,
    vocab_size: int,
    CL: int,
    ST: int,
    # Sampler config
    entropy_bound: float,
    # Tensor-parallel vocab sharding for the self-conditioning matmul.
    # ``embed_weight`` is vocab-sharded ([vocab/tp, hidden]) while ``probs``
    # spans the full vocab; [sc_vocab_start, sc_vocab_end) is this rank's slice.
    sc_vocab_start: int,
    sc_vocab_end: int,
    tp_size: int,
    tp_group_name: str,
    compute_sc: bool = True,
) -> None:
    """Compiled decode step: confidence → accept/renoise → convergence, as
    vectorized PyTorch ops over [num_decode, CL] tensors. The per-position
    statistics (argmax, Gumbel-max sample, entropy, softmax) come from one
    pass over the logits in ``sample_row_stats``."""
    num_decode = decode_slots.shape[0]
    device = decode_slots.device

    # ---- Phase 3: Confidence ----
    # A canvas truncated near max_model_len is zero-padded up to CL by the
    # caller; those padded rows are uniform (max entropy, argmax 0), so they
    # never trigger early convergence and are stable, and only the real
    # ``valid_canvas_len`` tokens are committed (num_sampled below).
    mean_entropy = token_entropy.mean(dim=-1)  # [num_decode]
    confident_tensor[decode_slots] = mean_entropy < confidence_threshold

    # ---- Phase 4: Entropy-bound acceptance mask ----
    sorted_ent, sorted_idx = torch.sort(token_entropy, dim=-1)
    cumsum_ent = torch.cumsum(sorted_ent, dim=-1)
    cummax_ent = torch.cummax(sorted_ent, dim=-1).values
    sorted_mask = (cumsum_ent - cummax_ent) <= entropy_bound
    eb_mask = torch.zeros_like(sorted_mask)
    eb_mask.scatter_(1, sorted_idx, sorted_mask)

    # ---- Phase 5: Post-sample ----
    is_commit = is_encoder_phase[decode_slots]  # [num_decode]
    is_denoise = ~is_commit
    cur_step = step_tensor[decode_slots].float()

    # Step update: +1 for denoise, reset to 0 for commit
    new_step_val = torch.where(
        is_denoise,
        (cur_step + 1).to(step_tensor.dtype),
        step_tensor.new_zeros(num_decode),
    )
    step_tensor[decode_slots] = new_step_val

    # Random tokens for renoise / canvas reinit
    # Renoise draws from the whole vocabulary even for a constrained read.
    # The model was trained on random-vocabulary noise. Noise drawn from the
    # allowed set looks like garbled text, and later steps denoise it into
    # more garbled text.
    random_tokens = torch.randint(
        0, vocab_size, (num_decode, CL), device=device, dtype=canvas.dtype
    )

    # Compute denoise canvas (accept/renoise). Pinned positions keep the seed.
    denoise_canvas = torch.where(eb_mask, new_tokens, random_tokens)
    denoise_canvas = torch.where(
        pin_mask[decode_slots], seed_canvas[decode_slots], denoise_canvas
    )

    # Canvas: commit → random reinit, denoise → accept/renoise result
    canvas[decode_slots] = torch.where(
        is_commit.unsqueeze(1), random_tokens, denoise_canvas
    )

    # History: write argmax_tokens for denoise requests at circular position
    hist_len = history_len_tensor[decode_slots]
    write_pos = hist_len % ST
    for i in range(ST):
        write_here = ((write_pos == i) & is_denoise).unsqueeze(1)
        history[decode_slots, i] = torch.where(
            write_here, argmax_tokens, history[decode_slots, i]
        )

    # Argmax canvas: update for denoise, preserve for commit
    argmax_canvas[decode_slots] = torch.where(
        is_denoise.unsqueeze(1), argmax_tokens, argmax_canvas[decode_slots]
    )

    # History length: increment for denoise, reset for commit
    new_hist_len = torch.where(is_denoise, hist_len + 1, hist_len.new_zeros(num_decode))
    history_len_tensor[decode_slots] = new_hist_len

    # ---- Phase 6: Stability + convergence ----
    ref = history[decode_slots, 0]
    mismatch = torch.zeros(num_decode, device=device, dtype=torch.int32)
    for h in range(1, ST):
        mismatch = mismatch + (ref != history[decode_slots, h]).sum(dim=-1).int()
    stable = mismatch == 0

    step_after = step_tensor[decode_slots]
    converged = (stable & confident_tensor[decode_slots] & (new_hist_len >= ST)) | (
        step_after >= max_steps_tensor[decode_slots]
    )
    # Commit done → denoise next (False); denoise converged → commit next (True)
    is_encoder_phase[decode_slots] = torch.where(
        is_commit, is_commit.new_zeros(num_decode), converged
    )

    emit = is_commit | (is_denoise & converged & read_only[decode_slots])
    sampled[decode_idx] = argmax_canvas[decode_slots].to(sampled.dtype) * emit[:, None]
    num_sampled[decode_idx] = (emit * valid_canvas_len).to(num_sampled.dtype)

    # SC soft embedding: store ``probs @ embed_weight`` (the value the next step's
    # self-conditioning MLP consumes) only for slots that will denoise next — i.e.
    # this step denoised AND it isn't about to commit (is_encoder_phase now False).
    # Masking here (rather than in the consumer) lets _apply_self_conditioning read
    # sc_embeds directly. Storing the [.., hidden] soft embed instead of the full
    # [.., vocab] probs avoids a giant persistent buffer.
    sc_keep = (is_denoise & ~is_encoder_phase[decode_slots])[:, None, None]
    if compute_sc:
        assert probs is not None
        # Self-conditioning soft embed = probs @ embed_tokens.weight. Under
        # tensor parallelism the embedding is vocab-sharded ([vocab/tp,
        # hidden]) while probs spans the full vocab, so each rank multiplies
        # its local vocab slice [sc_vocab_start, sc_vocab_end) and the
        # partials are summed across ranks.
        local_probs = probs[..., sc_vocab_start:sc_vocab_end].to(embed_weight.dtype)
        soft_embeds = torch.matmul(
            local_probs, embed_weight[: sc_vocab_end - sc_vocab_start]
        )
        if tp_size > 1:
            soft_embeds = torch.ops.vllm.all_reduce(
                soft_embeds, group_name=tp_group_name
            )
        soft_embeds = soft_embeds * normalizer
        # A pinned position holds its seed token as input. Zero its soft embed
        # too, or the model's own prediction there reaches the next step
        # through self-conditioning. The buffer is fp32 while the embedding is
        # in the model dtype: compiled code casts on the store but eager does
        # not, and the step runs eager once torch.compile hits its recompile
        # limit.
        sc_pin = (~pin_mask[decode_slots]).unsqueeze(-1)
        sc_embeds[decode_slots] = (soft_embeds * sc_keep * sc_pin).to(sc_embeds.dtype)
    else:
        # Every slot in this tile ends after this step, so the soft embed
        # would never be read. The matmul is a full pass over the vocabulary
        # matrix, so skip it.
        sc_embeds[decode_slots] = 0

    # Overwrite canvas with argmax for newly converged denoise requests
    newly_converged = (converged & is_denoise).unsqueeze(1)
    canvas[decode_slots] = torch.where(
        newly_converged, argmax_canvas[decode_slots], canvas[decode_slots]
    )
    is_encoder_phase[decode_slots] &= ~read_only[decode_slots]

    # ---- Phase 7: Copy canvas → draft_tokens for all slots ----
    draft_tokens[all_slots, :CL] = canvas[all_slots]

_concat_logprob_stashes(parts, cu_num_generated_tokens)

Join the logprobs stashed for the requests committing this step.

Each stash is as wide as the widest logprobs request in the batch at the step that request converged, so stashes from different steps can differ in width. Pad the narrow ones the way compute_topk_scores pads a mixed batch: token id 0 at -inf, which the output processor never reports.

Source code in vllm/model_executor/models/diffusion_gemma.py
def _concat_logprob_stashes(
    parts: list[LogprobsTensors], cu_num_generated_tokens: list[int]
) -> LogprobsTensors:
    """Join the logprobs stashed for the requests committing this step.

    Each stash is as wide as the widest logprobs request in the batch at the
    step that request converged, so stashes from different steps can differ
    in width. Pad the narrow ones the way compute_topk_scores pads a mixed
    batch: token id 0 at -inf, which the output processor never reports.
    """
    width = max(p.logprob_token_ids.shape[1] for p in parts)

    def pad(t: torch.Tensor, value: float) -> torch.Tensor:
        return F.pad(t, (0, width - t.shape[1]), value=value)

    return LogprobsTensors(
        logprob_token_ids=torch.cat([pad(p.logprob_token_ids, 0) for p in parts]),
        logprobs=torch.cat([pad(p.logprobs, float("-inf")) for p in parts]),
        selected_token_ranks=torch.cat([p.selected_token_ranks for p in parts]),
        cu_num_generated_tokens=cu_num_generated_tokens,
    )

_denoise_temperature(step_tensor, slots, max_denoising_steps, t_min, t_max)

The schedule's temperature for each slot at its current step.

Source code in vllm/model_executor/models/diffusion_gemma.py
@torch._dynamo.config.patch(recompile_limit=64)
@torch.compile(dynamic=True, backend=current_platform.simple_compile_backend)
def _denoise_temperature(
    step_tensor: torch.Tensor,
    slots: torch.Tensor,
    max_denoising_steps: float,
    t_min: float,
    t_max: float,
) -> torch.Tensor:
    """The schedule's temperature for each slot at its current step."""
    steps_f = step_tensor[slots].float()
    remaining = (max_denoising_steps - steps_f).clamp(min=1.0)
    return t_min + (t_max - t_min) * (remaining / max_denoising_steps)

_mask_rows_to_allowed(logits, row_starts, row_lens, allowed_per_row)

Mask every column outside a request's allowed ids on that request's rows. Request i owns rows [row_starts[i], +row_lens[i]); None leaves its rows alone. Returns a copy when any row is masked, so the caller's tensor (possibly the runner's) is never written.

Softmax over a masked row equals the K-space distribution the shared fast path computes, so both paths give the same reads. The mask value is a large finite negative rather than -inf: the entropy is probs times log-probs, and 0 * -inf is NaN, while 0 * -1e20 is 0. It stays finite after the greedy temperature clamp (1e-10) scales it by 1e10.

Source code in vllm/model_executor/models/diffusion_gemma.py
def _mask_rows_to_allowed(
    logits: torch.Tensor,
    row_starts: Sequence[int] | np.ndarray,
    row_lens: Sequence[int] | np.ndarray,
    allowed_per_row: Sequence[torch.Tensor | None],
) -> torch.Tensor:
    """Mask every column outside a request's allowed ids on that request's
    rows. Request i owns rows [row_starts[i], +row_lens[i]); None leaves its
    rows alone. Returns a copy when any row is masked, so the caller's tensor
    (possibly the runner's) is never written.

    Softmax over a masked row equals the K-space distribution the shared
    fast path computes, so both paths give the same reads. The mask value is
    a large finite negative rather than -inf: the entropy is probs times
    log-probs, and 0 * -inf is NaN, while 0 * -1e20 is 0. It stays finite
    after the greedy temperature clamp (1e-10) scales it by 1e10.
    """
    out: torch.Tensor | None = None
    keep = torch.zeros(logits.shape[-1], dtype=torch.bool, device=logits.device)
    for start, n, allowed in zip(row_starts, row_lens, allowed_per_row):
        if allowed is None:
            continue
        if out is None:
            out = logits.clone()
        keep.zero_()
        keep[allowed] = True
        out[int(start) : int(start) + int(n)].masked_fill_(~keep, _MASKED_LOGIT)
    return logits if out is None else out