Skip to content
Closed
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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -159,7 +159,7 @@ nixl-cu12 = false

[tool.uv.sources]
torch = { index = "pytorch-cu128" }
verifiers = { git = "https://github.com/PrimeIntellect-ai/verifiers.git", rev = "3b77145" }
verifiers = { git = "https://github.com/PrimeIntellect-ai/verifiers.git", rev = "350c1604e8a944cc9c84d0bae72fecba1fd62f1c" }
torchtitan = { git = "https://github.com/pytorch/torchtitan", rev = "a1fdd7e" }
dion = { git = "https://github.com/samsja/dion.git", rev = "d891eeb" }
transformers = { git = "https://github.com/huggingface/transformers.git", rev = "c1c3424" }
Expand Down
4 changes: 3 additions & 1 deletion src/prime_rl/configs/orchestrator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1075,7 +1075,7 @@ class OrchestratorConfig(BaseConfig):
Field(
description="Whether to use the renderer client (client-side tokenization via the ``renderers`` package, "
"served by ``/v1/generate``). Mutually exclusive with ``use_token_client``. When True, the "
"``[orchestrator.renderer]`` block (name / tool_parser / reasoning_parser / pool_size) "
"``[orchestrator.renderer]`` block (name / tool_parser / reasoning_parser / pool_size / keep_thinking) "
"applies; when False those fields must be left at their defaults. Not supported for VLMs — "
"VLMs must use the token client (TITO) so image preprocessing and chat templating stay server-side."
),
Expand Down Expand Up @@ -1207,6 +1207,8 @@ def validate_renderer_args(self):
renderer_args_set.append(f"renderer.reasoning_parser={self.renderer.reasoning_parser!r}")
if self.renderer.pool_size is not None:
renderer_args_set.append(f"renderer.pool_size={self.renderer.pool_size!r}")
if self.renderer.keep_thinking is not None:
renderer_args_set.append(f"renderer.keep_thinking={self.renderer.keep_thinking!r}")

if renderer_args_set:
raise ValueError(
Expand Down
11 changes: 11 additions & 0 deletions src/prime_rl/configs/shared.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,6 +186,17 @@ class RendererConfig(BaseConfig):
),
] = None

keep_thinking: Annotated[
bool | None,
Field(
description=(
"Override historical assistant reasoning block handling for renderers that support "
"keep_thinking. True preserves historical thinking, False disables it, and None "
"uses the selected renderer's default."
),
),
] = None


class ElasticConfig(BaseConfig):
"""Configures elastic inference pool with DNS-based service discovery.
Expand Down
2 changes: 2 additions & 0 deletions src/prime_rl/orchestrator/orchestrator.py
Original file line number Diff line number Diff line change
Expand Up @@ -925,6 +925,7 @@ async def setup_rollout_inference_pool(
renderer=config.renderer.name,
tool_parser=config.renderer.tool_parser,
reasoning_parser=config.renderer.reasoning_parser,
keep_thinking=config.renderer.keep_thinking,
)
logger.info(f"Initialized {type(renderer).__name__} for {config.model.name}")
inference_pool = await setup_inference_pool(
Expand All @@ -936,6 +937,7 @@ async def setup_rollout_inference_pool(
tool_parser=config.renderer.tool_parser,
reasoning_parser=config.renderer.reasoning_parser,
renderer_pool_size=config.renderer.pool_size,
renderer_keep_thinking=config.renderer.keep_thinking,
)
logger.info("Using direct renderer rollout client")
return renderer, inference_pool
Expand Down
7 changes: 7 additions & 0 deletions src/prime_rl/utils/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,7 @@ def __init__(
tool_parser: str | None = None,
reasoning_parser: str | None = None,
renderer_pool_size: int | None = None,
renderer_keep_thinking: bool | None = None,
):
renderer_model_name = model_name if train_client_type == "renderer" else None
self._train_clients = setup_clients(
Expand All @@ -78,6 +79,7 @@ def __init__(
tool_parser=tool_parser,
reasoning_parser=reasoning_parser,
renderer_pool_size=renderer_pool_size,
renderer_keep_thinking=renderer_keep_thinking,
)
self._eval_clients = setup_clients(client_config, client_type=eval_client_type)
self._admin_clients = setup_admin_clients(client_config)
Expand Down Expand Up @@ -129,6 +131,7 @@ async def setup_inference_pool(
tool_parser: str | None = None,
reasoning_parser: str | None = None,
renderer_pool_size: int | None = None,
renderer_keep_thinking: bool | None = None,
) -> InferencePool:
"""Create an inference pool from config (static or elastic)."""
logger = get_logger()
Expand All @@ -152,6 +155,7 @@ async def setup_inference_pool(
tool_parser=tool_parser,
reasoning_parser=reasoning_parser,
renderer_pool_size=renderer_pool_size,
renderer_keep_thinking=renderer_keep_thinking,
)

logger.info(
Expand All @@ -168,6 +172,7 @@ async def setup_inference_pool(
tool_parser=tool_parser,
reasoning_parser=reasoning_parser,
renderer_pool_size=renderer_pool_size,
renderer_keep_thinking=renderer_keep_thinking,
)


Expand All @@ -179,6 +184,7 @@ def setup_clients(
tool_parser: str | None = None,
reasoning_parser: str | None = None,
renderer_pool_size: int | None = None,
renderer_keep_thinking: bool | None = None,
) -> list[vf.ClientConfig]:
clients = []
client_idx = 0
Expand All @@ -194,6 +200,7 @@ def setup_clients(
renderer=renderer_name,
renderer_model_name=renderer_model_name,
renderer_pool_size=renderer_pool_size,
renderer_keep_thinking=renderer_keep_thinking,
tool_parser=tool_parser,
reasoning_parser=reasoning_parser,
api_base_url=base_url,
Expand Down
5 changes: 5 additions & 0 deletions src/prime_rl/utils/elastic.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,7 @@ def __init__(
tool_parser: str | None = None,
reasoning_parser: str | None = None,
renderer_pool_size: int | None = None,
renderer_keep_thinking: bool | None = None,
):
self.logger = get_logger()
self.client_config = client_config
Expand All @@ -125,6 +126,7 @@ def __init__(
self.tool_parser = tool_parser
self.reasoning_parser = reasoning_parser
self.renderer_pool_size = renderer_pool_size
self.renderer_keep_thinking = renderer_keep_thinking
self.router_url = client_config.router_url

self._servers: dict[str, ServerState] = {}
Expand Down Expand Up @@ -152,6 +154,7 @@ async def from_config(
tool_parser: str | None = None,
reasoning_parser: str | None = None,
renderer_pool_size: int | None = None,
renderer_keep_thinking: bool | None = None,
) -> ElasticInferencePool:
if client_config.elastic is None:
raise ValueError("Elastic inference pool requires elastic config")
Expand All @@ -164,6 +167,7 @@ async def from_config(
tool_parser=tool_parser,
reasoning_parser=reasoning_parser,
renderer_pool_size=renderer_pool_size,
renderer_keep_thinking=renderer_keep_thinking,
)
await pool.start()
return pool
Expand Down Expand Up @@ -214,6 +218,7 @@ def _rebuild_clients(self) -> None:
tool_parser=self.tool_parser,
reasoning_parser=self.reasoning_parser,
renderer_pool_size=self.renderer_pool_size,
renderer_keep_thinking=self.renderer_keep_thinking,
)
if urls
else []
Expand Down
3 changes: 3 additions & 0 deletions tests/unit/orchestrator/test_orchestrator_setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ async def run() -> None:
tool_parser=None,
reasoning_parser=None,
pool_size=None,
keep_thinking=True,
),
)
rollout_client_config = SimpleNamespace(base_url=["http://localhost:8000/v1"])
Expand Down Expand Up @@ -79,6 +80,7 @@ async def run() -> None:
renderer="qwen3_vl",
tool_parser=None,
reasoning_parser=None,
keep_thinking=True,
)
setup_pool_mock.assert_awaited_once_with(
rollout_client_config,
Expand All @@ -89,6 +91,7 @@ async def run() -> None:
tool_parser=None,
reasoning_parser=None,
renderer_pool_size=None,
renderer_keep_thinking=True,
)

asyncio.run(run())
21 changes: 21 additions & 0 deletions tests/unit/test_configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,27 @@ def test_removed_fused_lm_head_chunk_size_field_is_rejected():
TrainerModelConfig.model_validate({"fused_lm_head_chunk_size": "auto"})


@pytest.mark.parametrize("keep_thinking", [True, False])
def test_renderer_keep_thinking_requires_renderer_client(keep_thinking):
with pytest.raises(ValidationError, match="renderer.keep_thinking"):
OrchestratorConfig.model_validate({"renderer": {"keep_thinking": keep_thinking}})


def test_renderer_keep_thinking_is_accepted_with_renderer_client():
config = OrchestratorConfig.model_validate(
{
"use_token_client": False,
"use_renderer": True,
"renderer": {
"name": "nemotron3",
"keep_thinking": True,
},
}
)

assert config.renderer.keep_thinking is True


def test_selective_activation_checkpointing_requires_custom_impl():
with pytest.raises(ValidationError, match="Selective activation checkpointing requires model.impl='custom'"):
TrainerModelConfig.model_validate({"impl": "hf", "ac": {"mode": "selective"}})
18 changes: 18 additions & 0 deletions tests/unit/utils/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,7 @@ def test_setup_clients_assigns_renderer_and_dp_rank_headers():
assert [client.client_type for client in clients] == ["renderer", "renderer"]
assert [client.renderer for client in clients] == ["qwen3_vl", "qwen3_vl"]
assert [client.renderer_model_name for client in clients] == [None, None]
assert [client.renderer_keep_thinking for client in clients] == [None, None]
assert [client.api_base_url for client in clients] == ["http://worker-a:8000/v1"] * 2
assert [client.extra_headers["X-data-parallel-rank"] for client in clients] == ["0", "1"]
assert clients[0].extra_headers["X-Test"] == "test"
Expand All @@ -89,6 +90,22 @@ def test_setup_clients_assigns_renderer_model_name():
assert clients[0].renderer_model_name == "Qwen/Qwen3-VL-4B-Instruct"


def test_setup_clients_assigns_renderer_keep_thinking():
client_config = ClientConfig(
base_url=["http://worker-a:8000/v1"],
api_key_var="PRIME_API_KEY",
)

clients = setup_clients(
client_config,
client_type="renderer",
renderer_name="nemotron3",
renderer_keep_thinking=True,
)

assert clients[0].renderer_keep_thinking is True


def test_setup_clients_preserves_chat_client_defaults():
client_config = ClientConfig(
base_url=["http://worker-a:8000/v1"],
Expand All @@ -103,6 +120,7 @@ def test_setup_clients_preserves_chat_client_defaults():
client_type="openai_chat_completions",
renderer="auto",
renderer_model_name=None,
renderer_keep_thinking=None,
api_key_var="PRIME_API_KEY",
api_base_url="http://worker-a:8000/v1",
timeout=client_config.timeout,
Expand Down
2 changes: 2 additions & 0 deletions tests/unit/utils/test_elastic.py
Original file line number Diff line number Diff line change
Expand Up @@ -423,6 +423,7 @@ def test_elastic_clients_preserve_renderer_model_name_when_model_name_updates():
model_name="Qwen/Qwen3-VL-4B-Instruct",
train_client_type="renderer",
renderer_name="qwen3_vl",
renderer_keep_thinking=True,
)
pool._servers = {
"10.0.0.1": MagicMock(status="ready"),
Expand All @@ -437,6 +438,7 @@ def test_elastic_clients_preserve_renderer_model_name_when_model_name_updates():
client_type="renderer",
renderer="qwen3_vl",
renderer_model_name="Qwen/Qwen3-VL-4B-Instruct",
renderer_keep_thinking=True,
api_key_var="PRIME_API_KEY",
api_base_url="http://10.0.0.1:8000/v1",
timeout=1200,
Expand Down
6 changes: 3 additions & 3 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading