Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
82 changes: 47 additions & 35 deletions src/prime_rl/inference/vllm/serving_tokens.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,36 +2,35 @@

vLLM ships a generic tokens-in / tokens-out handler at
``vllm.entrypoints.scale_out.token_in_token_out.serving.ServingTokens`` that covers
prefix-cache salting, lora dispatch, multimodal features, prompt logprobs,
priority, ``data_parallel_rank`` header routing, server-side ``max_tokens``
defaulting and ``usage`` reporting. We subclass it for the bits still missing
from the upstream handler:
prefix-cache salting, lora dispatch, multimodal content parts and features,
prompt logprobs, priority, ``data_parallel_rank`` header routing, server-side
``max_tokens`` defaulting and ``usage`` reporting. We subclass it for the bits
still missing from the upstream handler:

1. Compact ``routed_experts`` export — when the engine emits routing
decisions, surface them as ``{data, shape, start, dtype}`` base64 raw-byte
objects (the form the PD router can merge and the renderers parse) instead
of upstream's single ``.npy`` base64 string.

2. ``kv_transfer_params`` bridging — upstream ``ServingTokens.serve_tokens``
parses ``request.kv_transfer_params`` but never threads it into the engine,
so PD disagg never fires on ``/inference/v1/generate``. Fixed upstream by
https://github.com/vllm-project/vllm/pull/42644, which missed the 0.28.0
cut — drop the bridge once we pin a release that includes it.
parses ``request.kv_transfer_params`` but never threads it into the engine.
Fixed upstream by https://github.com/vllm-project/vllm/pull/42644, which
missed the 0.28.0 cut.

Everything else (request/response schema, sampling params, error handling)
delegates to upstream so we track future vLLM changes for free.
3. Expanded prompt IDs — return the effective engine prompt after multimodal
placeholder expansion. Drop this once
https://github.com/vllm-project/vllm/pull/53187 is available in a release.

Everything else delegates to upstream so we track future vLLM changes for free.
"""

from __future__ import annotations

from collections.abc import AsyncGenerator
from collections.abc import AsyncGenerator, AsyncIterable
from typing import Any

from fastapi import Request
from vllm.entrypoints.openai.engine.protocol import (
ErrorResponse,
RequestResponseMetadata,
)
from vllm.entrypoints.openai.engine.protocol import ErrorResponse, RequestResponseMetadata
from vllm.entrypoints.scale_out.token_in_token_out.protocol import (
GenerateRequest,
GenerateResponse,
Expand All @@ -44,14 +43,14 @@


class PrimeRlGenerateResponseChoice(GenerateResponseChoice):
# Overrides upstream's base64 ``.npy`` string form with the compact
# ``{data, shape, start, dtype}`` object the PD router merges and the
# renderers parse.
# Overrides upstream's base64 ``.npy`` string form with the compact object
# the PD router merges and the renderers parse.
routed_experts: dict[str, Any] | None = None # type: ignore[assignment]


class PrimeRlGenerateResponse(GenerateResponse):
choices: list[PrimeRlGenerateResponseChoice]
prompt_token_ids: list[int] | None = None


class _GenerateRoutedExpertsCapture(RoutedExpertsCapture):
Expand All @@ -63,23 +62,29 @@ def post_process(self, response: GenerateResponse) -> PrimeRlGenerateResponse:
)
for choice in response.choices
]
return PrimeRlGenerateResponse(**{**dict(response), "choices": choices})
return PrimeRlGenerateResponse(**{**response.model_dump(exclude={"choices"}), "choices": choices})


class _PromptTokenIdsCapture:
def __init__(self, source: AsyncIterable[RequestOutput]) -> None:
self._source = source
self.prompt_token_ids: list[int] | None = None

async def __aiter__(self) -> AsyncGenerator[RequestOutput, None]:
async for output in self._source:
self.prompt_token_ids = output.prompt_token_ids
yield output


class PrimeRlServingTokens(ServingTokens):
"""ServingTokens + compact routed experts + PD kv_transfer_params bridging."""
"""ServingTokens with Prime's remaining response and PD extensions."""

async def serve_tokens(
self,
request: GenerateRequest,
raw_request: Request | None = None,
) -> GenerateResponse | ErrorResponse | AsyncGenerator[str, None]:
# Upstream parses ``request.kv_transfer_params`` but never threads it
# into the engine, so decode receives an empty NIXL handshake and
# re-prefills the prompt locally (~100x slower under concurrency).
# Bridge it through ``sampling_params.extra_args`` so the engine's KV
# connector picks the params up. Fixed upstream by vllm#42644 (merged
# after 0.28.0) — drop once we pin a release that includes it.
# Fixed upstream by vllm#42644; drop once it is included in the pin.
if request.kv_transfer_params is not None:
extra = request.sampling_params.extra_args or {}
extra["kv_transfer_params"] = request.kv_transfer_params
Expand All @@ -95,22 +100,29 @@ async def serve_tokens_full_generator( # type: ignore[override]
model_name: str,
request_metadata: RequestResponseMetadata,
) -> ErrorResponse | GenerateResponse:
# Capture routed_experts as vLLM streams request outputs, then post-process
# the final response into our GenerateResponse subclass so the encoded
# experts surface in the JSON.
capture: _GenerateRoutedExpertsCapture | None = None
routed_experts: _GenerateRoutedExpertsCapture | None = None
if self.model_config.enable_return_routed_experts:
capture = _GenerateRoutedExpertsCapture(
routed_experts = _GenerateRoutedExpertsCapture(
result_generator,
start=request.sampling_params.routed_experts_prompt_start,
)
result_generator = capture
result_generator = routed_experts

prompt_capture = _PromptTokenIdsCapture(result_generator)
response = await super().serve_tokens_full_generator(
request, result_generator, request_id, model_name, request_metadata
request,
prompt_capture,
request_id,
model_name,
request_metadata,
)

if capture is not None and isinstance(response, GenerateResponse):
response = capture.post_process(response)
if not isinstance(response, GenerateResponse):
return response

if routed_experts is not None:
response = routed_experts.post_process(response)
else:
response = PrimeRlGenerateResponse(**response.model_dump())
response.prompt_token_ids = prompt_capture.prompt_token_ids
return response
9 changes: 9 additions & 0 deletions src/prime_rl/multimodal/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
from prime_rl.multimodal.base import ForwardPolicy, MaterializedMM, MultimodalAdapter
from prime_rl.multimodal.registry import get_multimodal_adapter

__all__ = [
"ForwardPolicy",
"MaterializedMM",
"MultimodalAdapter",
"get_multimodal_adapter",
]
40 changes: 40 additions & 0 deletions src/prime_rl/multimodal/base.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import Any, Protocol

import torch
from PIL.Image import Image


@dataclass(frozen=True)
class ForwardPolicy:
pass_position_ids: bool = True
requires_mm_token_type_ids: bool = False
defer_context_parallelism: bool = False


@dataclass(frozen=True)
class MaterializedMM:
kwargs: dict[str, torch.Tensor]
forward_policy: ForwardPolicy


class MultimodalAdapter(Protocol):
model_types: frozenset[str]
forward_policy: ForwardPolicy

def materialize(
self,
image_processor: Any,
images: list[Image],
placeholder_lengths: list[int],
) -> MaterializedMM: ...


def required_tensors(values: Any, keys: tuple[str, ...]) -> dict[str, torch.Tensor]:
data = dict(values)
missing = [key for key in keys if key not in data]
if missing:
raise ValueError(f"Image processor did not return {', '.join(missing)}")
return {key: torch.as_tensor(data[key]).contiguous() for key in keys}
33 changes: 33 additions & 0 deletions src/prime_rl/multimodal/kimi_k25.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
from __future__ import annotations

from typing import Any

from PIL.Image import Image

from prime_rl.multimodal.base import ForwardPolicy, MaterializedMM, required_tensors


class KimiK25Adapter:
model_types = frozenset({"kimi_k25"})
forward_policy = ForwardPolicy()

def materialize(
self,
image_processor: Any,
images: list[Image],
placeholder_lengths: list[int],
) -> MaterializedMM:
preprocess = getattr(image_processor, "preprocess", None)
if preprocess is None:
raise ValueError("Kimi image processor is missing preprocess")
media = [{"type": "image", "image": image} for image in images]
kwargs = required_tensors(
preprocess(media, return_tensors="pt"),
("pixel_values", "grid_thws"),
)
lengths = [1] * len(kwargs["grid_thws"].reshape(-1, 3))
if lengths != placeholder_lengths:
raise ValueError(
f"Kimi image placeholder lengths differ from vLLM: expected {placeholder_lengths}, got {lengths}"
)
return MaterializedMM(kwargs=kwargs, forward_policy=self.forward_policy)
35 changes: 35 additions & 0 deletions src/prime_rl/multimodal/qwen_vl.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
from __future__ import annotations

from typing import Any

from PIL.Image import Image

from prime_rl.multimodal.base import ForwardPolicy, MaterializedMM, required_tensors


class QwenVLAdapter:
model_types = frozenset({"qwen3_vl", "qwen3_vl_moe", "qwen3_5", "qwen3_5_moe"})
forward_policy = ForwardPolicy(
pass_position_ids=False,
requires_mm_token_type_ids=True,
defer_context_parallelism=True,
)

def materialize(
self,
image_processor: Any,
images: list[Image],
placeholder_lengths: list[int],
) -> MaterializedMM:
kwargs = required_tensors(
image_processor(images=images, return_tensors="pt"),
("pixel_values", "image_grid_thw"),
)
merge_size = int(image_processor.merge_size)
# HF Qwen-VL / renderer pad count: T*H*W / merge_size^2.
lengths = [int(grid.prod()) // (merge_size * merge_size) for grid in kwargs["image_grid_thw"].reshape(-1, 3)]
if lengths != placeholder_lengths:
raise ValueError(
f"Qwen image placeholder lengths differ from vLLM: expected {placeholder_lengths}, got {lengths}"
)
return MaterializedMM(kwargs=kwargs, forward_policy=self.forward_policy)
17 changes: 17 additions & 0 deletions src/prime_rl/multimodal/registry.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
from __future__ import annotations

from prime_rl.multimodal.base import MultimodalAdapter
from prime_rl.multimodal.kimi_k25 import KimiK25Adapter
from prime_rl.multimodal.qwen_vl import QwenVLAdapter

_ADAPTERS = (QwenVLAdapter(), KimiK25Adapter())
_BY_MODEL_TYPE: dict[str, MultimodalAdapter] = {
model_type: adapter for adapter in _ADAPTERS for model_type in adapter.model_types
}


def get_multimodal_adapter(model_type: str) -> MultimodalAdapter:
try:
return _BY_MODEL_TYPE[model_type]
except KeyError as exc:
raise NotImplementedError(f"Raw image training is not implemented for model type {model_type!r}") from exc

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Missing KL mismatch validation table for new models

Medium Severity

This PR introduces new custom multimodal adapters (QwenVLAdapter for qwen3_vl/qwen3_vl_moe/qwen3_5/qwen3_5_moe and KimiK25Adapter for kimi_k25) with distinct ForwardPolicy configurations that change how position_ids, mm_token_type_ids, and context parallelism behave during the model forward pass. Per project rules, any PR introducing a new custom model must include a table showing mean KL mismatch across 20 steps on a math environment with batch_size=64, with all entries below 0.015. No such table is present in the PR description.

Additional Locations (2)
Fix in Cursor Fix in Web

Triggered by project rule: BugBot Instructions

Reviewed by Cursor Bugbot for commit ab4ca01. Configure here.

Loading