Skip to content

vllm.v1.worker.gpu.spec_decode.dflash.cudagraph

Classes:

DFlashCudaGraphManager

Bases: CudaGraphManager

DFlash CudaGraphManager for the parallel-drafting query forward, building its own attention metadata from scratch.

Methods:

  • capture –

    precompute_context_kv(num_reqs) is captured ahead of the query

Source code in vllm/v1/worker/gpu/spec_decode/dflash/cudagraph.py
class DFlashCudaGraphManager(CudaGraphManager):
    """DFlash CudaGraphManager for the parallel-drafting query forward,
    building its own attention metadata from scratch."""

    def capture(
        self,
        forward_fn: Callable,
        input_buffers: InputBuffers,
        block_tables: BlockTables,
        attn_groups: list[list[AttentionGroup]],
        kv_cache_config: KVCacheConfig,
        max_model_len: int,
        causal: bool | Mapping[int, bool],
        precompute_context_kv: Callable[[int], None],
        progress_bar_desc: str = "Capturing CUDA graphs",
    ) -> None:
        """``precompute_context_kv(num_reqs)`` is captured ahead of the query
        forward."""

        def create_forward_fn(
            desc: BatchExecutionDescriptor,
            warmup: bool,
        ) -> Callable[[CUDAGraphMode], None]:
            num_tokens = desc.num_tokens
            num_reqs = desc.num_reqs or min(num_tokens, self.max_num_reqs)
            num_tokens_across_dp = (
                torch.full((self.dp_size,), num_tokens, dtype=torch.int32, device="cpu")
                if self.dp_size > 1
                else None
            )
            attn_state = _prepare_dflash_inputs_to_capture(
                num_reqs,
                num_tokens,
                input_buffers,
                block_tables,
                attn_groups,
                kv_cache_config,
                max_model_len,
                skip_attn=(desc.cg_mode == CUDAGraphMode.PIECEWISE),
                causal=causal,
            )
            attn_metadata, slot_mappings = attn_state

            def forward(cg_mode: CUDAGraphMode) -> None:
                precompute_context_kv(num_reqs)
                forward_fn(
                    num_reqs,
                    num_tokens,
                    attn_metadata,
                    slot_mappings,
                    num_tokens_across_dp,
                    cg_mode,
                )

            return forward

        super().capture(create_forward_fn, progress_bar_desc)

capture(forward_fn, input_buffers, block_tables, attn_groups, kv_cache_config, max_model_len, causal, precompute_context_kv, progress_bar_desc='Capturing CUDA graphs')

precompute_context_kv(num_reqs) is captured ahead of the query forward.

Source code in vllm/v1/worker/gpu/spec_decode/dflash/cudagraph.py
def capture(
    self,
    forward_fn: Callable,
    input_buffers: InputBuffers,
    block_tables: BlockTables,
    attn_groups: list[list[AttentionGroup]],
    kv_cache_config: KVCacheConfig,
    max_model_len: int,
    causal: bool | Mapping[int, bool],
    precompute_context_kv: Callable[[int], None],
    progress_bar_desc: str = "Capturing CUDA graphs",
) -> None:
    """``precompute_context_kv(num_reqs)`` is captured ahead of the query
    forward."""

    def create_forward_fn(
        desc: BatchExecutionDescriptor,
        warmup: bool,
    ) -> Callable[[CUDAGraphMode], None]:
        num_tokens = desc.num_tokens
        num_reqs = desc.num_reqs or min(num_tokens, self.max_num_reqs)
        num_tokens_across_dp = (
            torch.full((self.dp_size,), num_tokens, dtype=torch.int32, device="cpu")
            if self.dp_size > 1
            else None
        )
        attn_state = _prepare_dflash_inputs_to_capture(
            num_reqs,
            num_tokens,
            input_buffers,
            block_tables,
            attn_groups,
            kv_cache_config,
            max_model_len,
            skip_attn=(desc.cg_mode == CUDAGraphMode.PIECEWISE),
            causal=causal,
        )
        attn_metadata, slot_mappings = attn_state

        def forward(cg_mode: CUDAGraphMode) -> None:
            precompute_context_kv(num_reqs)
            forward_fn(
                num_reqs,
                num_tokens,
                attn_metadata,
                slot_mappings,
                num_tokens_across_dp,
                cg_mode,
            )

        return forward

    super().capture(create_forward_fn, progress_bar_desc)