Skip to content

vllm.v1.worker.gpu.sample.watermark

Functions:

_repeated_context_mask_cpu(all_token_ids, req_indices, prompt_lens, total_lens, contexts, max_history=None, include_prompt=False, skip_partial_context=False, history_offsets=None)

Reference implementation; the parity tests check the Triton kernel against it.

Source code in vllm/v1/worker/gpu/sample/watermark.py
def _repeated_context_mask_cpu(
    all_token_ids: torch.Tensor,
    req_indices: torch.Tensor,
    prompt_lens: torch.Tensor,
    total_lens: torch.Tensor,
    contexts: torch.Tensor,
    max_history: int | None = None,
    include_prompt: bool = False,
    skip_partial_context: bool = False,
    history_offsets: torch.Tensor | None = None,
) -> torch.Tensor:
    """Reference implementation; the parity tests check the Triton kernel against it."""
    repeated = torch.zeros(len(req_indices), dtype=torch.bool)
    for row, req_idx_tensor in enumerate(req_indices):
        req_idx = int(req_idx_tensor)
        if req_idx < 0:
            continue
        prompt_len = int(prompt_lens[req_idx])
        total_len = int(total_lens[req_idx])
        if skip_partial_context and total_len - prompt_len < contexts.shape[-1]:
            repeated[row] = True
            continue
        sequence_start = 0 if include_prompt else prompt_len
        history_tokens = all_token_ids[req_idx, sequence_start:total_len].tolist()
        prefix = [-1] * contexts.shape[-1]
        current_context = tuple(contexts[row].tolist())
        history_offset = 0 if history_offsets is None else int(history_offsets[row])
        history_start = (
            max(0, len(history_tokens) - max(0, max_history - history_offset))
            if max_history is not None
            else 0
        )
        for history_pos, token_id in enumerate(history_tokens):
            if history_pos >= history_start and (
                tuple(prefix[-contexts.shape[-1] :]) == current_context
            ):
                repeated[row] = True
                break
            prefix.append(token_id)
    return repeated

philox_gumbel_block_argmax(logits, mask, block_idx, contexts_row_ptr, key_0, key_1, CONTEXT_WIDTH, BLOCK_SIZE)

Apply the detector-compatible PRF to one vocabulary block.

Source code in vllm/v1/worker/gpu/sample/watermark.py
@triton.jit
def philox_gumbel_block_argmax(
    logits,
    mask,
    block_idx,
    contexts_row_ptr,
    key_0,
    key_1,
    CONTEXT_WIDTH: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """Apply the detector-compatible PRF to one vocabulary block."""
    tl.static_assert(BLOCK_SIZE % 4 == 0)
    state_0, state_1, state_2, state_3 = _philox_context_state(
        contexts_row_ptr, key_0, key_1, CONTEXT_WIDTH
    )
    groups = block_idx * (BLOCK_SIZE // 4) + tl.arange(0, BLOCK_SIZE // 4)
    output_0, output_1, output_2, output_3 = _philox_candidate_words(
        groups, state_0, state_1, state_2, state_3, key_0, key_1
    )
    # Lay the per-group words out in token order: 4g+0, 4g+1, 4g+2, 4g+3.
    words = tl.interleave(
        tl.interleave(output_0, output_2),
        tl.interleave(output_1, output_3),
    )
    values = tl.where(
        mask,
        _philox_gumbel_from_logits(logits.to(tl.float32), words),
        float("-inf"),
    )
    return tl.max(values, axis=0, return_indices=True)

repeated_context_mask(all_token_ids, req_indices, prompt_lens, total_lens, contexts, max_history=None, include_prompt=False, skip_partial_context=False, history_offsets=None, local_positions=None, num_speculative_steps=0)

Return, per row, whether the row's context already occurred in its history.

Parameters:

  • all_token_ids

    (Tensor) –

    [max_num_reqs, max_model_len] token ids of every request.

  • req_indices

    (Tensor) –

    Request slot per sampled row; -1 marks a padding row, which is reported as not repeated.

  • prompt_lens

    (Tensor) –

    Prompt length per request slot.

  • total_lens

    (Tensor) –

    Prompt plus generated length per request slot.

  • contexts

    (Tensor) –

    [num_rows, context_width] context per row, padded with -1 before the start of the scanned history.

  • max_history

    (int | None, default: None ) –

    Number of most recent history positions searched, or None for all of them. The window compared at each position reaches context_width tokens further back.

  • include_prompt

    (bool, default: False ) –

    Search the prompt as well as the generated tokens.

  • skip_partial_context

    (bool, default: False ) –

    Mark contexts containing start padding so they use ordinary sampling.

  • history_offsets

    (Tensor | None, default: None ) –

    Number of newer, non-committed positions preceding each context. These positions count toward max_history. Must be 1-D over rows; a stride-0 broadcast row is fine.

  • local_positions

    (Tensor | None, default: None ) –

    Position within each request's speculative block. When provided, compare each context with earlier contexts in that block.

  • num_speculative_steps

    (int, default: 0 ) –

    Maximum number of draft tokens in the block.

Source code in vllm/v1/worker/gpu/sample/watermark.py
def repeated_context_mask(
    all_token_ids: torch.Tensor,
    req_indices: torch.Tensor,
    prompt_lens: torch.Tensor,
    total_lens: torch.Tensor,
    contexts: torch.Tensor,
    max_history: int | None = None,
    include_prompt: bool = False,
    skip_partial_context: bool = False,
    history_offsets: torch.Tensor | None = None,
    local_positions: torch.Tensor | None = None,
    num_speculative_steps: int = 0,
) -> torch.Tensor:
    """Return, per row, whether the row's context already occurred in its history.

    Args:
        all_token_ids: `[max_num_reqs, max_model_len]` token ids of every request.
        req_indices: Request slot per sampled row; -1 marks a padding row, which
            is reported as not repeated.
        prompt_lens: Prompt length per request slot.
        total_lens: Prompt plus generated length per request slot.
        contexts: `[num_rows, context_width]` context per row, padded with -1
            before the start of the scanned history.
        max_history: Number of most recent history positions searched, or
            ``None`` for all of them. The window compared at each position
            reaches `context_width` tokens further back.
        include_prompt: Search the prompt as well as the generated tokens.
        skip_partial_context: Mark contexts containing start padding so they use
            ordinary sampling.
        history_offsets: Number of newer, non-committed positions preceding each
            context. These positions count toward `max_history`. Must be 1-D
            over rows; a stride-0 broadcast row is fine.
        local_positions: Position within each request's speculative block. When
            provided, compare each context with earlier contexts in that block.
        num_speculative_steps: Maximum number of draft tokens in the block.

    """
    if max_history is not None and max_history < 1:
        raise ValueError("max_history must be positive or None")
    if (local_positions is None) != (num_speculative_steps == 0):
        raise ValueError(
            "local_positions and positive num_speculative_steps must be "
            "provided together"
        )
    if all_token_ids.device.type == "cpu":
        repeated = _repeated_context_mask_cpu(
            all_token_ids,
            req_indices,
            prompt_lens,
            total_lens,
            contexts,
            max_history,
            include_prompt,
            skip_partial_context,
            history_offsets,
        )
        if local_positions is not None:
            max_offset = num_speculative_steps + 1
            if max_history is not None:
                max_offset = min(max_offset, max_history + 1)
            for offset in range(1, max_offset):
                prior_contexts = torch.cat(
                    (contexts[:offset], contexts[:-offset]), dim=0
                )
                prior_requests = torch.cat(
                    (req_indices[:offset], req_indices[:-offset]), dim=0
                )
                prior_local_pos = torch.cat(
                    (local_positions[:offset], local_positions[:-offset]), dim=0
                )
                repeated |= (
                    (req_indices >= 0)
                    & (local_positions >= offset)
                    & (req_indices == prior_requests)
                    & (local_positions == prior_local_pos + offset)
                    & (contexts == prior_contexts).all(dim=-1)
                )
            if include_prompt:
                repeated |= (contexts < 0).any(dim=-1)
        return repeated

    if contexts.stride(-1) != 1:
        contexts = contexts.contiguous()
    # Only history_offsets passes its row stride to the kernel.
    req_indices = req_indices.contiguous()
    if local_positions is not None:
        local_positions = local_positions.contiguous()
    # `prompt_lens` and `total_lens` are indexed flat by request slot, and
    # `all_token_ids` takes a row stride but is read flat within the row. Those
    # three are contiguous by contract; every caller owns them outright.
    repeated = torch.empty(len(req_indices), dtype=torch.bool, device=contexts.device)
    _repeated_context_mask_kernel[(len(req_indices),)](
        repeated,
        all_token_ids,
        all_token_ids.stride(0),
        req_indices,
        prompt_lens,
        total_lens,
        history_offsets,
        0 if history_offsets is None else history_offsets.stride(0),
        contexts,
        contexts.stride(0),
        local_positions,
        None,
        0,
        0,
        None,
        0,
        None,
        CONTEXT_WIDTH=contexts.shape[-1],
        NUM_SPECULATIVE_STEPS=num_speculative_steps,
        MAX_HISTORY=0 if max_history is None else max_history,
        INCLUDE_PROMPT=include_prompt,
        SKIP_PARTIAL_CONTEXT=skip_partial_context,
        HAS_HISTORY_OFFSETS=history_offsets is not None,
        SPECULATIVE_CONTEXTS=local_positions is not None,
        DRAFT_CONTEXTS=False,
        BLOCK=512,
    )
    return repeated