Skip to content

vllm.v1.pool.flash_maxsim

Fused Triton kernels for late-interaction (MaxSim) scoring.

Modules:

  • flash_maxsim_rerank –

    Flash-MaxSim rerank: one query vs variable-length docs in a packed tensor.

Functions:

flash_maxsim_rerank_direct(Q, batch_tensor, doc_offsets, doc_lengths, max_seqlen_d)

TRUE zero-copy: score query against docs scattered in a batch tensor.

The kernel reads doc embeddings directly from batch_tensor at the positions specified by doc_offsets. No torch.stack, no torch.cat, no copy of any kind. The batch tensor is the model's output.

Memory for doc scoring: 0 bytes additional.

Parameters:

  • Q

    (Tensor) –

    [Lq, d] — single query embedding (from cache)

  • batch_tensor

    (Tensor) –

    [total_tokens, d] — the model's projected output tensor. Contains ALL requests' tokens (queries + docs + others).

  • doc_offsets

    (Tensor) –

    [B] int32 — start token index of each doc in batch_tensor

  • doc_lengths

    (Tensor) –

    [B] int32 — number of tokens per doc

  • max_seqlen_d

    (int) –

    int — max(doc_lengths)

Returns:

  • scores ( Tensor ) –

    [B] float32 — one MaxSim score per document

Source code in vllm/v1/pool/flash_maxsim/flash_maxsim_rerank.py
def flash_maxsim_rerank_direct(
    Q: torch.Tensor,
    batch_tensor: torch.Tensor,
    doc_offsets: torch.Tensor,
    doc_lengths: torch.Tensor,
    max_seqlen_d: int,
) -> torch.Tensor:
    """TRUE zero-copy: score query against docs scattered in a batch tensor.

    The kernel reads doc embeddings directly from batch_tensor at the
    positions specified by doc_offsets. No torch.stack, no torch.cat,
    no copy of any kind. The batch tensor is the model's output.

    Memory for doc scoring: 0 bytes additional.

    Args:
        Q: [Lq, d] — single query embedding (from cache)
        batch_tensor: [total_tokens, d] — the model's projected output tensor.
            Contains ALL requests' tokens (queries + docs + others).
        doc_offsets: [B] int32 — start token index of each doc in batch_tensor
        doc_lengths: [B] int32 — number of tokens per doc
        max_seqlen_d: int — max(doc_lengths)

    Returns:
        scores: [B] float32 — one MaxSim score per document

    """
    assert Q.dim() == 2, f"Q must be 2D [Lq, d], got {Q.dim()}D"
    assert batch_tensor.dim() == 2, (
        f"batch_tensor must be 2D, got {batch_tensor.dim()}D"
    )
    assert Q.shape[1] == batch_tensor.shape[1]

    return _run_rerank_kernel(Q, batch_tensor, doc_offsets, doc_lengths, max_seqlen_d)