Skip to content

vllm.model_executor.layers.indexer_topk

Top-k kernels for the DSA sparse attention indexer.

Classes:

Functions:

SparseIndexerTopk

Bases: Module

The sparse indexer's decode top-k stage.

Selects among the available top-k kernels (see kernel_config.sparse_indexer_topk_backend) and runs the chosen one.

Methods:

  • forward –

    Run the resolved decode top-k implementation, writing into

  • resolve_backend –

    Resolve the decode top-k implementation from the configured

Source code in vllm/model_executor/layers/indexer_topk.py
class SparseIndexerTopk(torch.nn.Module):
    """The sparse indexer's decode top-k stage.

    Selects among the available top-k kernels (see
    kernel_config.sparse_indexer_topk_backend) and runs the chosen one.
    """

    def __init__(self, backend: str | None = None) -> None:
        super().__init__()
        if backend is None:
            backend = (
                get_current_vllm_config().kernel_config.sparse_indexer_topk_backend
            )
        self._backend = backend
        self._is_cuda = current_platform.is_cuda()
        self._has_deep_select = self._is_cuda and (
            current_platform.is_device_capability_family(100)
        )
        self._has_flashinfer_topk = has_flashinfer()
        self._cooperative_capable = self._is_cuda and (
            current_platform.has_device_capability(90)
            and not current_platform.is_device_capability_family(120)
        )
        self._is_rocm = current_platform.is_rocm()
        self._aiter_enabled = False
        self._aiter_capable = False
        if self._is_rocm:
            from vllm._aiter_ops import rocm_aiter_ops

            self._aiter_enabled = bool(rocm_aiter_ops.is_enabled())
            self._aiter_capable = bool(rocm_aiter_ops.is_indexer_top_k_enabled())

    def resolve_backend(
        self, logits: torch.Tensor, topk_tokens: int, num_rows: int
    ) -> str:
        """Resolve the decode top-k implementation from the configured
        backend ("auto" = the pre-existing chain, or a validated explicit
        value)."""
        if self._backend == "auto":
            return self._resolve_auto(logits, topk_tokens, num_rows)

        failures: list[str] = []
        if self._backend == "cooperative":
            failures = self._cooperative_constraints(logits, topk_tokens, num_rows)
        elif self._backend == "persistent":
            if not self._is_cuda:
                failures.append("requires a CUDA platform")
            if topk_tokens not in (512, 1024, 2048):
                failures.append(
                    f"topk_tokens must be in (512, 1024, 2048), got {topk_tokens}"
                )
        elif self._backend == "deep_select":
            if not self._is_cuda:
                failures.append("requires a CUDA platform")
            elif not self._has_deep_select:
                failures.append("requires SM100a/SM103a (10.x device family)")
            elif not is_deep_select_supported(logits, topk_tokens):
                failures.append(
                    f"inputs violate DeepSelect's constraints: dtype={logits.dtype},"
                    f" stride={logits.stride()}, topk={topk_tokens}"
                )
        elif self._backend == "aiter":
            if not self._is_rocm:
                failures.append("requires a ROCm platform")
            elif not self._aiter_enabled:
                failures.append("requires AITER (VLLM_ROCM_USE_AITER=1)")
            elif not self._aiter_capable:
                failures.append(
                    "AITER has no tuned indexer top-k kernels for this GPU arch"
                )
            if topk_tokens not in (512, 1024, 2048, 4096):
                failures.append(
                    f"topk_tokens must be in (512, 1024, 2048, 4096), got {topk_tokens}"
                )
        elif self._backend == "flashinfer":
            if not self._is_cuda:
                failures.append("requires a CUDA platform")
            if not self._has_flashinfer_topk:
                failures.append(
                    "flashinfer.topk.top_k_ragged_transform is not importable"
                )
            if logits.dtype != torch.float32 or logits.stride(1) != 1:
                failures.append(
                    f"requires fp32 logits with stride(1) == 1, got "
                    f"dtype={logits.dtype}, stride={logits.stride()}"
                )
        if failures:
            raise RuntimeError(
                f"sparse_indexer_topk_backend='{self._backend}' was requested, but: "
                + "; ".join(failures)
            )
        return self._backend

    def _resolve_auto(
        self, logits: torch.Tensor, topk_tokens: int, num_rows: int
    ) -> str:
        """The priority chain: cooperative -> persistent -> per_row.
        aiter/deep_select/flashinfer/torch are opt-in only.

        cooperative_topk is preferred whenever it is applicable, i.e. within
        its AUTO_COOPERATIVE_MAX_ROWS row limit; larger batches go to
        persistent_topk.

        aiter not auto-resolved as it does not outperform the per-row on every shape,
        hence we delegate to the caller to decide when to use it.
        """
        if not self._cooperative_constraints(logits, topk_tokens, num_rows):
            return "cooperative"
        if self._is_cuda and topk_tokens in (512, 1024, 2048):
            return "persistent"
        return "per_row"

    def _cooperative_constraints(
        self, logits: torch.Tensor, topk_tokens: int, num_rows: int
    ) -> list[str]:
        """Unmet constraints of cooperative_topk (empty when applicable)."""
        failures = []
        if not self._is_cuda:
            failures.append("requires a CUDA platform")
        if topk_tokens not in (512, 1024, 2048):
            failures.append(
                f"topk_tokens must be in (512, 1024, 2048), got {topk_tokens}"
            )
        if num_rows > AUTO_COOPERATIVE_MAX_ROWS:
            failures.append(
                f"num_rows must be <= {AUTO_COOPERATIVE_MAX_ROWS}, got {num_rows}"
            )
        if logits.stride(0) % 4 != 0:
            failures.append(
                f"logits.stride(0) must be divisible by 4, got {logits.stride(0)}"
            )
        if self._is_cuda and not self._cooperative_capable:
            failures.append("requires SM90+ and is not supported on the SM12x family")
        return failures

    @staticmethod
    def _row_ends(seq_lens: torch.Tensor, next_n: int, num_rows: int) -> torch.Tensor:
        """Per-row exclusive end offsets (int32, (num_rows,)) for top-k
        kernels that take ragged lengths (DeepSelect, FlashInfer, torch
        reference).

        seq_lens is (B, next_n) per-row effective lens for native spec decode
        and (B, 1) otherwise, in which case per-row lens are derived the same
        way as the other decode top-k kernels.
        """
        if seq_lens.numel() == num_rows:
            row_ends = seq_lens.reshape(-1)
        else:
            next_n_offsets = torch.arange(
                next_n, dtype=torch.int32, device=seq_lens.device
            )
            row_ends = (
                (seq_lens.reshape(-1, 1) - next_n + 1 + next_n_offsets)
                .clamp_(min=0)
                .reshape(-1)
            )
        assert row_ends.dtype == torch.int32
        return row_ends

    def forward(
        self,
        logits: torch.Tensor,
        seq_lens: torch.Tensor,
        next_n: int,
        topk_indices: torch.Tensor,
        topk_tokens: int,
        max_seq_len: int,
    ) -> None:
        """Run the resolved decode top-k implementation, writing into
        topk_indices (int32, -1 fill for rows shorter than topk_tokens).
        """
        backend = self.resolve_backend(logits, topk_tokens, logits.shape[0])
        if backend == "deep_select":
            row_ends = self._row_ends(seq_lens, next_n, logits.shape[0])
            deep_select_topk(logits, topk_tokens, end=row_ends, output_idx=topk_indices)
        elif backend == "cooperative":
            (topk_workspace,) = current_workspace_manager().get_simultaneous(
                ((RADIX_TOPK_WORKSPACE_SIZE,), torch.uint8),
            )
            torch.ops._C.cooperative_topk(
                logits,
                seq_lens,
                topk_indices,
                topk_workspace,
                topk_tokens,
                max_seq_len,
            )
        elif backend == "persistent":
            (topk_workspace,) = current_workspace_manager().get_simultaneous(
                ((RADIX_TOPK_WORKSPACE_SIZE,), torch.uint8),
            )
            # persistent_topk's last argument is the host-side "max seq_len
            # across rows" bound (it gates the sampled_topk path and clamps
            # per-row lengths), not the padded row width: logits is sized to
            # max_model_len, so passing its width would incorrectly enable the
            # long-row kernels on short-context batches.
            torch.ops._C.persistent_topk(
                logits,
                seq_lens,
                topk_indices,
                topk_workspace,
                topk_tokens,
                max_seq_len,
            )
        elif backend == "flashinfer":
            # Deferred: importing flashinfer initializes CUDA at import time.
            from flashinfer.topk import top_k_ragged_transform

            row_ends = self._row_ends(seq_lens, next_n, logits.shape[0])
            offsets = torch.zeros(
                logits.shape[0], dtype=torch.int32, device=logits.device
            )
            # top_k_ragged_transform selects within [0, row_ends[i]) per
            # row and -1-fills past the row length.
            indices = top_k_ragged_transform(logits, offsets, row_ends, topk_tokens)
            topk_indices.copy_(indices)
        elif backend == "torch":
            # Debug reference: mask everything past each row's end, then topk.
            row_ends = self._row_ends(seq_lens, next_n, logits.shape[0])
            cols = torch.arange(logits.shape[1], device=logits.device)
            masked = logits.masked_fill(
                cols.unsqueeze(0) >= row_ends.unsqueeze(1), float("-inf")
            )
            indices = masked.topk(topk_tokens, dim=-1).indices
            in_range = torch.arange(topk_tokens, device=logits.device).unsqueeze(
                0
            ) < row_ends.unsqueeze(1)
            indices = torch.where(in_range, indices, -1)
            topk_indices.copy_(indices)
        elif backend == "aiter":
            from vllm._aiter_ops import rocm_aiter_ops

            # AITER's decode kernel only implements the 1D/next_n == 1 form,
            # which is equivalent once the ragged ends are expanded.
            rocm_aiter_ops.indexer_top_k_decode(
                logits,
                1,
                self._row_ends(seq_lens, next_n, logits.shape[0]),
                topk_indices,
                topk_tokens,
            )
        else:
            assert backend == "per_row", f"unknown topk backend: {backend}"
            self._per_row(logits, seq_lens, next_n, topk_indices, topk_tokens)

    @staticmethod
    def _per_row(
        logits: torch.Tensor,
        seq_lens: torch.Tensor,
        next_n: int,
        topk_indices: torch.Tensor,
        topk_tokens: int,
    ) -> None:
        ops.top_k_per_row_decode(
            logits,
            next_n,
            seq_lens,
            topk_indices,
            logits.shape[0],
            logits.stride(0),
            logits.stride(1),
            topk_tokens,
        )

_cooperative_constraints(logits, topk_tokens, num_rows)

Unmet constraints of cooperative_topk (empty when applicable).

Source code in vllm/model_executor/layers/indexer_topk.py
def _cooperative_constraints(
    self, logits: torch.Tensor, topk_tokens: int, num_rows: int
) -> list[str]:
    """Unmet constraints of cooperative_topk (empty when applicable)."""
    failures = []
    if not self._is_cuda:
        failures.append("requires a CUDA platform")
    if topk_tokens not in (512, 1024, 2048):
        failures.append(
            f"topk_tokens must be in (512, 1024, 2048), got {topk_tokens}"
        )
    if num_rows > AUTO_COOPERATIVE_MAX_ROWS:
        failures.append(
            f"num_rows must be <= {AUTO_COOPERATIVE_MAX_ROWS}, got {num_rows}"
        )
    if logits.stride(0) % 4 != 0:
        failures.append(
            f"logits.stride(0) must be divisible by 4, got {logits.stride(0)}"
        )
    if self._is_cuda and not self._cooperative_capable:
        failures.append("requires SM90+ and is not supported on the SM12x family")
    return failures

_resolve_auto(logits, topk_tokens, num_rows)

The priority chain: cooperative -> persistent -> per_row. aiter/deep_select/flashinfer/torch are opt-in only.

cooperative_topk is preferred whenever it is applicable, i.e. within its AUTO_COOPERATIVE_MAX_ROWS row limit; larger batches go to persistent_topk.

aiter not auto-resolved as it does not outperform the per-row on every shape, hence we delegate to the caller to decide when to use it.

Source code in vllm/model_executor/layers/indexer_topk.py
def _resolve_auto(
    self, logits: torch.Tensor, topk_tokens: int, num_rows: int
) -> str:
    """The priority chain: cooperative -> persistent -> per_row.
    aiter/deep_select/flashinfer/torch are opt-in only.

    cooperative_topk is preferred whenever it is applicable, i.e. within
    its AUTO_COOPERATIVE_MAX_ROWS row limit; larger batches go to
    persistent_topk.

    aiter not auto-resolved as it does not outperform the per-row on every shape,
    hence we delegate to the caller to decide when to use it.
    """
    if not self._cooperative_constraints(logits, topk_tokens, num_rows):
        return "cooperative"
    if self._is_cuda and topk_tokens in (512, 1024, 2048):
        return "persistent"
    return "per_row"

_row_ends(seq_lens, next_n, num_rows) staticmethod

Per-row exclusive end offsets (int32, (num_rows,)) for top-k kernels that take ragged lengths (DeepSelect, FlashInfer, torch reference).

seq_lens is (B, next_n) per-row effective lens for native spec decode and (B, 1) otherwise, in which case per-row lens are derived the same way as the other decode top-k kernels.

Source code in vllm/model_executor/layers/indexer_topk.py
@staticmethod
def _row_ends(seq_lens: torch.Tensor, next_n: int, num_rows: int) -> torch.Tensor:
    """Per-row exclusive end offsets (int32, (num_rows,)) for top-k
    kernels that take ragged lengths (DeepSelect, FlashInfer, torch
    reference).

    seq_lens is (B, next_n) per-row effective lens for native spec decode
    and (B, 1) otherwise, in which case per-row lens are derived the same
    way as the other decode top-k kernels.
    """
    if seq_lens.numel() == num_rows:
        row_ends = seq_lens.reshape(-1)
    else:
        next_n_offsets = torch.arange(
            next_n, dtype=torch.int32, device=seq_lens.device
        )
        row_ends = (
            (seq_lens.reshape(-1, 1) - next_n + 1 + next_n_offsets)
            .clamp_(min=0)
            .reshape(-1)
        )
    assert row_ends.dtype == torch.int32
    return row_ends

forward(logits, seq_lens, next_n, topk_indices, topk_tokens, max_seq_len)

Run the resolved decode top-k implementation, writing into topk_indices (int32, -1 fill for rows shorter than topk_tokens).

Source code in vllm/model_executor/layers/indexer_topk.py
def forward(
    self,
    logits: torch.Tensor,
    seq_lens: torch.Tensor,
    next_n: int,
    topk_indices: torch.Tensor,
    topk_tokens: int,
    max_seq_len: int,
) -> None:
    """Run the resolved decode top-k implementation, writing into
    topk_indices (int32, -1 fill for rows shorter than topk_tokens).
    """
    backend = self.resolve_backend(logits, topk_tokens, logits.shape[0])
    if backend == "deep_select":
        row_ends = self._row_ends(seq_lens, next_n, logits.shape[0])
        deep_select_topk(logits, topk_tokens, end=row_ends, output_idx=topk_indices)
    elif backend == "cooperative":
        (topk_workspace,) = current_workspace_manager().get_simultaneous(
            ((RADIX_TOPK_WORKSPACE_SIZE,), torch.uint8),
        )
        torch.ops._C.cooperative_topk(
            logits,
            seq_lens,
            topk_indices,
            topk_workspace,
            topk_tokens,
            max_seq_len,
        )
    elif backend == "persistent":
        (topk_workspace,) = current_workspace_manager().get_simultaneous(
            ((RADIX_TOPK_WORKSPACE_SIZE,), torch.uint8),
        )
        # persistent_topk's last argument is the host-side "max seq_len
        # across rows" bound (it gates the sampled_topk path and clamps
        # per-row lengths), not the padded row width: logits is sized to
        # max_model_len, so passing its width would incorrectly enable the
        # long-row kernels on short-context batches.
        torch.ops._C.persistent_topk(
            logits,
            seq_lens,
            topk_indices,
            topk_workspace,
            topk_tokens,
            max_seq_len,
        )
    elif backend == "flashinfer":
        # Deferred: importing flashinfer initializes CUDA at import time.
        from flashinfer.topk import top_k_ragged_transform

        row_ends = self._row_ends(seq_lens, next_n, logits.shape[0])
        offsets = torch.zeros(
            logits.shape[0], dtype=torch.int32, device=logits.device
        )
        # top_k_ragged_transform selects within [0, row_ends[i]) per
        # row and -1-fills past the row length.
        indices = top_k_ragged_transform(logits, offsets, row_ends, topk_tokens)
        topk_indices.copy_(indices)
    elif backend == "torch":
        # Debug reference: mask everything past each row's end, then topk.
        row_ends = self._row_ends(seq_lens, next_n, logits.shape[0])
        cols = torch.arange(logits.shape[1], device=logits.device)
        masked = logits.masked_fill(
            cols.unsqueeze(0) >= row_ends.unsqueeze(1), float("-inf")
        )
        indices = masked.topk(topk_tokens, dim=-1).indices
        in_range = torch.arange(topk_tokens, device=logits.device).unsqueeze(
            0
        ) < row_ends.unsqueeze(1)
        indices = torch.where(in_range, indices, -1)
        topk_indices.copy_(indices)
    elif backend == "aiter":
        from vllm._aiter_ops import rocm_aiter_ops

        # AITER's decode kernel only implements the 1D/next_n == 1 form,
        # which is equivalent once the ragged ends are expanded.
        rocm_aiter_ops.indexer_top_k_decode(
            logits,
            1,
            self._row_ends(seq_lens, next_n, logits.shape[0]),
            topk_indices,
            topk_tokens,
        )
    else:
        assert backend == "per_row", f"unknown topk backend: {backend}"
        self._per_row(logits, seq_lens, next_n, topk_indices, topk_tokens)

resolve_backend(logits, topk_tokens, num_rows)

Resolve the decode top-k implementation from the configured backend ("auto" = the pre-existing chain, or a validated explicit value).

Source code in vllm/model_executor/layers/indexer_topk.py
def resolve_backend(
    self, logits: torch.Tensor, topk_tokens: int, num_rows: int
) -> str:
    """Resolve the decode top-k implementation from the configured
    backend ("auto" = the pre-existing chain, or a validated explicit
    value)."""
    if self._backend == "auto":
        return self._resolve_auto(logits, topk_tokens, num_rows)

    failures: list[str] = []
    if self._backend == "cooperative":
        failures = self._cooperative_constraints(logits, topk_tokens, num_rows)
    elif self._backend == "persistent":
        if not self._is_cuda:
            failures.append("requires a CUDA platform")
        if topk_tokens not in (512, 1024, 2048):
            failures.append(
                f"topk_tokens must be in (512, 1024, 2048), got {topk_tokens}"
            )
    elif self._backend == "deep_select":
        if not self._is_cuda:
            failures.append("requires a CUDA platform")
        elif not self._has_deep_select:
            failures.append("requires SM100a/SM103a (10.x device family)")
        elif not is_deep_select_supported(logits, topk_tokens):
            failures.append(
                f"inputs violate DeepSelect's constraints: dtype={logits.dtype},"
                f" stride={logits.stride()}, topk={topk_tokens}"
            )
    elif self._backend == "aiter":
        if not self._is_rocm:
            failures.append("requires a ROCm platform")
        elif not self._aiter_enabled:
            failures.append("requires AITER (VLLM_ROCM_USE_AITER=1)")
        elif not self._aiter_capable:
            failures.append(
                "AITER has no tuned indexer top-k kernels for this GPU arch"
            )
        if topk_tokens not in (512, 1024, 2048, 4096):
            failures.append(
                f"topk_tokens must be in (512, 1024, 2048, 4096), got {topk_tokens}"
            )
    elif self._backend == "flashinfer":
        if not self._is_cuda:
            failures.append("requires a CUDA platform")
        if not self._has_flashinfer_topk:
            failures.append(
                "flashinfer.topk.top_k_ragged_transform is not importable"
            )
        if logits.dtype != torch.float32 or logits.stride(1) != 1:
            failures.append(
                f"requires fp32 logits with stride(1) == 1, got "
                f"dtype={logits.dtype}, stride={logits.stride()}"
            )
    if failures:
        raise RuntimeError(
            f"sparse_indexer_topk_backend='{self._backend}' was requested, but: "
            + "; ".join(failures)
        )
    return self._backend

_get_empty_and_aligned_tensor(dim0, dim1, device, dtype)

Tensor with shape (dim0, dim1) whose stride(0) is 32B-aligned.

Source code in vllm/model_executor/layers/indexer_topk.py
def _get_empty_and_aligned_tensor(
    dim0: int, dim1: int, device: torch.device, dtype: torch.dtype
) -> torch.Tensor:
    """Tensor with shape (dim0, dim1) whose stride(0) is 32B-aligned."""
    output_stride_requirement = get_deep_select_stride_requirement()[1] // (
        dtype.itemsize
    )
    assert output_stride_requirement > 0
    dim1_rounded = (
        (dim1 + output_stride_requirement - 1) // output_stride_requirement
    ) * output_stride_requirement
    return torch.empty((dim0, dim1_rounded), device=device, dtype=dtype)[:, :dim1]

deep_select_topk(input, topk, end=None, output_idx=None, indices_dtype=torch.int32)

Select the top-k indices per row of input with DeepSelect.

Parameters:

  • input

    (Tensor) –

    (num_rows, vocab_size), bf16 or fp32. stride(1) must be 1 and stride(0) must be 1024B-aligned.

  • topk

    (int) –

    Number of elements to select per row; must be <= 4096.

  • end

    (Tensor | None, default: None ) –

    Optional (num_rows,) int32 tensor with the exclusive right boundary of each row. Rows with end[i] < topk get their remaining indices filled with -1.

  • output_idx

    (Tensor | None, default: None ) –

    Optional preallocated (num_rows, topk) output tensor whose stride(0) is 32B-aligned (e.g. a slice of a wider buffer).

  • indices_dtype

    (dtype, default: int32 ) –

    Output dtype when output_idx is not provided.

Returns:

  • Tensor –

    The (num_rows, topk) indices tensor.

Source code in vllm/model_executor/layers/indexer_topk.py
def deep_select_topk(
    input: torch.Tensor,
    topk: int,
    end: torch.Tensor | None = None,
    output_idx: torch.Tensor | None = None,
    indices_dtype: torch.dtype = torch.int32,
) -> torch.Tensor:
    """Select the top-k indices per row of `input` with DeepSelect.

    Args:
        input: (num_rows, vocab_size), bf16 or fp32. stride(1) must be 1 and
            stride(0) must be 1024B-aligned.
        topk: Number of elements to select per row; must be <= 4096.
        end: Optional (num_rows,) int32 tensor with the exclusive right
            boundary of each row. Rows with `end[i] < topk` get their
            remaining indices filled with -1.
        output_idx: Optional preallocated (num_rows, topk) output tensor whose
            stride(0) is 32B-aligned (e.g. a slice of a wider buffer).
        indices_dtype: Output dtype when `output_idx` is not provided.

    Returns:
        The (num_rows, topk) indices tensor.

    """
    assert input.dim() == 2 and input.stride(1) == 1
    assert (
        input.stride(0) * input.element_size() % get_deep_select_stride_requirement()[0]
        == 0
    )

    num_rows = input.shape[0]
    if output_idx is None:
        output_idx = _get_empty_and_aligned_tensor(
            num_rows, topk, input.device, indices_dtype
        )
    else:
        assert output_idx.dtype == indices_dtype
        assert output_idx.shape[0] >= num_rows and output_idx.shape[1] >= topk

    torch.ops.deep_select.topk(
        input,
        topk,
        None,  # begin is not supported
        end,
        False,  # sorted_value
        False,  # sorted_index
        None,  # output_value
        output_idx,
        None,  # output_idx_offset
        IDX_OOB_FILL_VALUE,
        float("-inf"),  # value_oob_fill_value
        False,  # return_value
        False,  # abort_when_nan_found: dummy/capture passes can feed NaN
        # logits from uninitialized KV cache; the reference topk kernels do
        # not trap on NaN, and trapping would abort CUDA graph capture.
    )
    return output_idx

get_deep_select_stride_requirement() cached

Stride alignment requirement (input, output) in bytes.

Source code in vllm/model_executor/layers/indexer_topk.py
@functools.lru_cache(maxsize=1)
def get_deep_select_stride_requirement() -> tuple[int, int]:
    """Stride alignment requirement (input, output) in bytes."""
    return torch.ops.deep_select.get_alignment_requirement()

is_deep_select_supported(input, topk)

Whether the kernel accepts this input (dtype/stride/topk constraints).

Source code in vllm/model_executor/layers/indexer_topk.py
def is_deep_select_supported(input: torch.Tensor, topk: int) -> bool:
    """Whether the kernel accepts this input (dtype/stride/topk constraints)."""
    return (
        topk <= 4096
        and input.shape[1] < 2**23
        and input.dtype in (torch.float32, torch.bfloat16)
        and input.stride(1) == 1
        and input.stride(0)
        * input.element_size()
        % get_deep_select_stride_requirement()[0]
        == 0
    )