-
Notifications
You must be signed in to change notification settings - Fork 418
Train from inline multimodal rollouts #3320
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from 12 commits
cb20c6a
27cb537
ef12808
59bbc3d
09f280f
538f732
6b5de34
4418fb1
2f89b88
e96b231
c1b7a7d
2e3a0c3
d27af25
ab4ca01
2ad1838
577123e
0bd7b84
133fa23
7c0737c
af60d8a
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| +23 −15 | README.md | |
| +2 −7 | docs/renderer-config.md | |
| +21 −5 | pyproject.toml | |
| +11 −15 | renderers/__init__.py | |
| +85 −181 | renderers/base.py | |
| +75 −48 | renderers/client.py | |
| +2 −3 | renderers/deepseek_v3.py | |
| +2 −3 | renderers/default.py | |
| +5 −6 | renderers/gemma4.py | |
| +2 −3 | renderers/glm45.py | |
| +2 −3 | renderers/glm5.py | |
| +2 −3 | renderers/gpt_oss.py | |
| +2 −3 | renderers/hy3.py | |
| +5 −6 | renderers/inkling.py | |
| +2 −3 | renderers/kimi_k2.py | |
| +30 −14 | renderers/kimi_k25.py | |
| +2 −3 | renderers/laguna_s21.py | |
| +4 −5 | renderers/laguna_xs2.py | |
| +2 −3 | renderers/llama_3.py | |
| +2 −3 | renderers/minimax_m2.py | |
| +2 −3 | renderers/nemotron3.py | |
| +3 −4 | renderers/prime_qwen3.py | |
| +2 −3 | renderers/qwen3.py | |
| +50 −32 | renderers/qwen35.py | |
| +50 −32 | renderers/qwen3_vl.py | |
| +1 −0 | tests/conftest.py | |
| +1 −0 | tests/test_bridge.py | |
| +46 −1 | tests/test_client.py | |
| +4 −0 | tests/test_disabled_thinking_stability.py | |
| +18 −12 | tests/test_qwen38.py | |
| +0 −26 | tests/test_renderer_config.py | |
| +1 −0 | tests/test_renderer_config_parity.py | |
| +1 −0 | tests/test_roundtrip.py | |
| +14 −1 | uv.lock |
| +1 −1 | AGENTS.md | |
| +1 −1 | pyproject.toml | |
| +69 −0 | skills/release/SKILL.md | |
| +16 −0 | tests/v1/test_e2e.py | |
| +55 −0 | tests/v1/test_graph.py | |
| +66 −0 | tests/v1/test_trace.py | |
| +2 −1 | verifiers/v1/__init__.py | |
| +42 −8 | verifiers/v1/acp/__init__.py | |
| +45 −11 | verifiers/v1/acp/runner.py | |
| +11 −4 | verifiers/v1/agent.py | |
| +29 −26 | verifiers/v1/clients/train.py | |
| +2 −6 | verifiers/v1/dialects/responses.py | |
| +4 −39 | verifiers/v1/errors.py | |
| +78 −150 | verifiers/v1/graph.py | |
| +6 −0 | verifiers/v1/harnesses/__init__.py | |
| +40 −2 | verifiers/v1/harnesses/null/program.py | |
| +6 −0 | verifiers/v1/harnesses/prime_agent/__init__.py | |
| +311 −0 | verifiers/v1/harnesses/prime_agent/harness.py | |
| +50 −48 | verifiers/v1/harnesses/rlm/harness.py | |
| +1 −0 | verifiers/v1/harnesses/terminus_2/program.py | |
| +175 −86 | verifiers/v1/interception/server.py | |
| +19 −26 | verifiers/v1/rollout.py | |
| +11 −2 | verifiers/v1/runtimes/base.py | |
| +24 −11 | verifiers/v1/session.py | |
| +9 −25 | verifiers/v1/trace.py | |
| +1 −4 | verifiers/v1/types.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", | ||
| ] |
| 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} |
| 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) |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,34 @@ | ||
| 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) | ||
| 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) |
| 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 | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Missing KL mismatch validation table for new modelsMedium Severity This PR introduces new custom multimodal adapters ( Additional Locations (2)Triggered by project rule: BugBot Instructions Reviewed by Cursor Bugbot for commit ab4ca01. Configure here. |
||


Uh oh!
There was an error while loading. Please reload this page.