Skip to content

vllm.v1.worker.gpu.mm.rope

Classes:

  • RopeState

    State for multi-dimensional (M-RoPE) positions.

Functions:

  • get_rope_state

    Create a RopeState if the model uses multi-dimensional RoPE.

RopeState

State for multi-dimensional (M-RoPE) positions.

num_dims is the number of position channels the model consumes, which the model config derives from its M-RoPE sections.

NOTE: positions is implemented with one additional dummy position on purpose to make it non-contiguous so that it can work with torch compile. See detailed explanation in https://github.com/vllm-project/vllm/pull/12128#discussion_r1926431923

NOTE: For text-only inputs, each dimension has identical position IDs, making M-RoPE functionally equivalent to 1D-RoPE. See page 5 of https://arxiv.org/abs/2409.12191

Methods:

Source code in vllm/v1/worker/gpu/mm/rope.py
class RopeState:
    """State for multi-dimensional (M-RoPE) positions.

    `num_dims` is the number of position channels the model consumes, which
    the model config derives from its M-RoPE sections.

    NOTE: `positions` is implemented with one additional dummy position on
    purpose to make it non-contiguous so that it can work with torch compile.
    See detailed explanation in
    https://github.com/vllm-project/vllm/pull/12128#discussion_r1926431923

    NOTE: For text-only inputs, each dimension has identical position IDs,
    making M-RoPE functionally equivalent to 1D-RoPE.
    See page 5 of https://arxiv.org/abs/2409.12191
    """

    def __init__(
        self,
        num_dims: int,
        max_num_reqs: int,
        max_num_tokens: int,
        max_model_len: int,
        device: torch.device,
    ):
        self.num_dims = num_dims
        self.max_num_reqs = max_num_reqs
        self.max_num_tokens = max_num_tokens
        self.max_model_len = max_model_len
        self.device = device

        # NOTE(woosuk): This tensor can be extremely large (e.g., several GBs)
        # wasting a lot of CPU memory.
        self.prefill_positions = StagedWriteTensor(
            (max_num_reqs * num_dims, max_model_len),
            dtype=torch.int32,
            device=device,
            uva_instead_of_gpu=True,
        )
        self.positions = torch.zeros(
            (num_dims, max_num_tokens + 1), dtype=torch.int64, device=device
        )

        self.prefill_delta = UvaBackedTensor(max_num_reqs, dtype=torch.int32)

    def init_prefill_positions(
        self,
        req_idx: int,
        model: nn.Module,
        prefill_token_ids: list[int],
        mm_features: list,
    ) -> None:
        mrope_model = cast(SupportsMRoPE, model)
        prefill_positions, delta = mrope_model.get_mrope_input_positions(
            prefill_token_ids, mm_features
        )
        self.prefill_delta.np[req_idx] = delta

        for i in range(self.num_dims):
            pos = prefill_positions[i].tolist()
            self.prefill_positions.stage_write(self.num_dims * req_idx + i, 0, pos)

    def apply_staged_writes(self) -> None:
        self.prefill_positions.apply_write()
        self.prefill_delta.copy_to_uva()

    def get_positions(self, num_tokens: int) -> torch.Tensor:
        return self.positions[:, :num_tokens]

    def read_prefill_positions(self, req_idx: int, length: int) -> torch.Tensor:
        """Return staged per-request prefill positions as [num_dims, length]."""
        base = self.num_dims * req_idx
        return self.prefill_positions.gpu[base : base + self.num_dims, :length]

    def update_prefill_positions(
        self, req_idx: int, positions: torch.Tensor, delta: int
    ) -> None:
        """Overwrite a request's staged prefill positions with recomputed values."""
        base = self.num_dims * req_idx
        length = positions.shape[1]
        self.prefill_positions.gpu[base : base + self.num_dims, :length].copy_(
            positions
        )
        self.prefill_delta.np[req_idx] = delta

    def prepare_positions(
        self,
        idx_mapping: torch.Tensor,
        query_start_loc: torch.Tensor,
        prefill_lens: torch.Tensor,
        num_computed_tokens: torch.Tensor,
    ) -> None:
        num_reqs = idx_mapping.shape[0]
        _prepare_rope_positions_kernel[(num_reqs,)](
            self.positions,
            self.positions.stride(0),
            self.prefill_positions.gpu,
            self.num_dims * self.max_model_len,
            self.max_model_len,
            self.prefill_delta.gpu,
            idx_mapping,
            query_start_loc,
            prefill_lens,
            num_computed_tokens,
            BLOCK_SIZE=1024,
            NUM_DIMS=self.num_dims,
        )

read_prefill_positions(req_idx, length)

Return staged per-request prefill positions as [num_dims, length].

Source code in vllm/v1/worker/gpu/mm/rope.py
def read_prefill_positions(self, req_idx: int, length: int) -> torch.Tensor:
    """Return staged per-request prefill positions as [num_dims, length]."""
    base = self.num_dims * req_idx
    return self.prefill_positions.gpu[base : base + self.num_dims, :length]

update_prefill_positions(req_idx, positions, delta)

Overwrite a request's staged prefill positions with recomputed values.

Source code in vllm/v1/worker/gpu/mm/rope.py
def update_prefill_positions(
    self, req_idx: int, positions: torch.Tensor, delta: int
) -> None:
    """Overwrite a request's staged prefill positions with recomputed values."""
    base = self.num_dims * req_idx
    length = positions.shape[1]
    self.prefill_positions.gpu[base : base + self.num_dims, :length].copy_(
        positions
    )
    self.prefill_delta.np[req_idx] = delta

get_rope_state(model_config, model, max_num_reqs, max_num_tokens, max_model_len, device)

Create a RopeState if the model uses multi-dimensional RoPE.

Source code in vllm/v1/worker/gpu/mm/rope.py
def get_rope_state(
    model_config: ModelConfig,
    model: nn.Module,
    max_num_reqs: int,
    max_num_tokens: int,
    max_model_len: int,
    device: torch.device,
) -> RopeState | None:
    """Create a RopeState if the model uses multi-dimensional RoPE."""
    if not model_config.uses_mrope:
        return None

    assert isinstance(model, SupportsMRoPE)
    return RopeState(
        num_dims=model_config.mrope_num_dims,
        max_num_reqs=max_num_reqs,
        max_num_tokens=max_num_tokens,
        max_model_len=max_model_len,
        device=device,
    )