Skip to content

vllm_omni.diffusion.attention.backends.flashinfer_attn

HAS_FLASHINFER module-attribute

HAS_FLASHINFER = True

logger module-attribute

logger = init_logger(__name__)

FlashInferAttentionBackend

Bases: AttentionBackend

accept_output_buffer class-attribute instance-attribute

accept_output_buffer: bool = True

get_impl_cls staticmethod

get_impl_cls() -> type[FlashInferAttentionImpl]

get_name staticmethod

get_name() -> str

get_supported_head_sizes staticmethod

get_supported_head_sizes() -> list[int]

supports_attention_mask classmethod

supports_attention_mask(
    attention_spec: object | None = None,
) -> bool

FlashInferAttentionImpl

Bases: AttentionImpl

backend_explicit instance-attribute

backend_explicit = bool(
    extra_impl_args.get("backend_explicit", False)
)

causal instance-attribute

causal = causal

device instance-attribute

device = torch.device(
    "cuda", torch.accelerator.current_device_index()
)

dtype_qk instance-attribute

dtype_qk = self._check_dtype(
    quant.get("dtype_qk"), "dtype_qk", self._QK_DTYPES
)

dtype_vo instance-attribute

dtype_vo = self._check_dtype(
    quant.get("dtype_vo"), "dtype_vo", self._VO_DTYPES
)

flashinfer_backend instance-attribute

flashinfer_backend = self._select_backend(
    requested_backend, device=self.device
)

softmax_scale instance-attribute

softmax_scale = softmax_scale

forward_cuda

forward_cuda(
    query: Tensor,
    key: Tensor,
    value: Tensor,
    attn_metadata: AttentionMetadata | None = None,
) -> Tensor