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
7 changes: 4 additions & 3 deletions src/prime_rl/orchestrator/eval_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,13 +7,15 @@
import verifiers as vf

from prime_rl.orchestrator.config import EvalSamplingConfig
from prime_rl.orchestrator.utils import apply_client_sampling_overrides
from prime_rl.orchestrator.vf_utils import evaluate, get_completion_len
from prime_rl.utils.config import ClientConfig
from prime_rl.utils.logger import get_logger
from prime_rl.utils.monitor import get_monitor
from prime_rl.utils.utils import capitalize


def get_eval_sampling_args(sampling_config: EvalSamplingConfig) -> dict[str, Any]:
def get_eval_sampling_args(sampling_config: EvalSamplingConfig, client_config: ClientConfig) -> dict[str, Any]:
"""Get sampling args for evaluation."""
# Initialize sampling args
sampling_args: dict[str, Any] = {}
Expand Down Expand Up @@ -41,8 +43,7 @@ def get_eval_sampling_args(sampling_config: EvalSamplingConfig) -> dict[str, Any
extra_body["repetition_penalty"] = sampling_config.repetition_penalty

sampling_args["extra_body"] = extra_body

return sampling_args
return apply_client_sampling_overrides(sampling_args, client_config)


def _pass_at_k(n: int, c: int, k: int) -> float:
Expand Down
4 changes: 2 additions & 2 deletions src/prime_rl/orchestrator/orchestrator.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,7 +199,7 @@ async def orchestrate(config: OrchestratorConfig):
env_ids = [strip_env_version(env.id) for env in config.eval.env]
eval_envs = [vf.load_environment(env_id, **env.args) for env_id, env in zip(env_ids, config.eval.env)]
eval_env_names = [env.name or env_id for env_id, env in zip(env_ids, config.eval.env)]
eval_sampling_args = get_eval_sampling_args(config.eval.sampling)
eval_sampling_args = get_eval_sampling_args(config.eval.sampling, config.client)
eval_env_addresses = []

for env_id, env, eval_env_name in zip(env_ids, config.eval.env, eval_env_names):
Expand Down Expand Up @@ -422,7 +422,7 @@ async def orchestrate(config: OrchestratorConfig):

# Schedule generating the training batch
temperature = compute_temperature(progress.step, config.sampling, config.max_steps)
sampling_args = get_sampling_args(config.sampling, temperature=temperature)
sampling_args = get_sampling_args(config.sampling, temperature=temperature, client_config=config.client)
scheduler.set_sampling_args(sampling_args)
train_task = asyncio.create_task(scheduler.generate_batch(step=progress.step))

Expand Down
2 changes: 1 addition & 1 deletion src/prime_rl/orchestrator/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ def __init__(
self.strict_async_level = strict_async_level
self.lora_name = lora_name
initial_temp = compute_temperature(step=0, sampling_config=config.sampling, max_steps=config.max_steps)
self.sampling_args = get_sampling_args(config.sampling, temperature=initial_temp)
self.sampling_args = get_sampling_args(config.sampling, temperature=initial_temp, client_config=config.client)
self.model_name = self.config.model.name
self.json_logging = config.log.json_logging

Expand Down
27 changes: 23 additions & 4 deletions src/prime_rl/orchestrator/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@

from prime_rl.orchestrator.config import SamplingConfig
from prime_rl.transport import TrainingSample
from prime_rl.utils.config import ClientConfig
from prime_rl.utils.utils import (
format_num,
format_time,
Expand All @@ -37,22 +38,40 @@ async def get_semaphore() -> AsyncContextManager:
return SEMAPHORE


def get_sampling_args(sampling_config: SamplingConfig, temperature: float) -> dict:
def apply_client_sampling_overrides(
sampling_args: dict[str, Any],
client_config: ClientConfig,
) -> dict[str, Any]:
sampling_args = dict(sampling_args)
extra_body = dict(sampling_args.get("extra_body") or {})
extra_body.update(client_config.extra_body_overrides)
sampling_args["extra_body"] = extra_body

sampling_args.update({k: v for k, v in client_config.sampling_overrides.items() if k != "extra_body"})
return sampling_args


def get_sampling_args(sampling_config: SamplingConfig, temperature: float, client_config: ClientConfig) -> dict:
# Convert SamplingConfig to vLLM OAI sampling args
# https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html#extra-parameters_2
sampling_args = dict(sampling_config)
sampling_args.pop("temp_scheduler", None)
sampling_args["temperature"] = temperature
sampling_args["top_p"] = 1.0
sampling_args["logprobs"] = True
sampling_args["extra_body"] = {
**sampling_config.extra_body,
"return_token_ids": True, # Always return token IDs
"top_k": -1,
"min_p": 0.0,
**sampling_config.extra_body,
}
sampling_args["extra_body"]["min_tokens"] = sampling_args.pop("min_tokens")
sampling_args["extra_body"]["repetition_penalty"] = sampling_args.pop("repetition_penalty")

sampling_args = apply_client_sampling_overrides(sampling_args, client_config)

# Token-level outputs are required by rollout-to-training conversion.
sampling_args["logprobs"] = True
sampling_args["extra_body"]["return_token_ids"] = True

return sampling_args


Expand Down
4 changes: 2 additions & 2 deletions src/prime_rl/utils/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,7 @@ async def setup_inference_pool(client_config: ClientConfig, model_name: str) ->

logger.info(
f"Initializing static inference pool (base_url={', '.join(client_config.base_url)}, "
f"api_key_var={client_config.api_key_var}, headers={client_config.headers})"
f"api_key_var={client_config.api_key_var}, client_type={client_config.client_type}, headers={client_config.headers})"
)
return StaticInferencePool(
clients=setup_clients(client_config),
Expand All @@ -120,7 +120,7 @@ def setup_clients(client_config: ClientConfig) -> list[vf.ClientConfig]:
def setup_client(client_idx: int, base_url: str) -> vf.ClientConfig:
return vf.ClientConfig(
client_idx=client_idx,
client_type="openai_chat_completions_token",
client_type=client_config.client_type,
api_base_url=base_url,
api_key_var=client_config.api_key_var,
timeout=client_config.timeout,
Expand Down
32 changes: 31 additions & 1 deletion src/prime_rl/utils/config.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import Annotated, Literal
from typing import Annotated, Any, Literal

from pydantic import Field

Expand All @@ -19,6 +19,7 @@ class ModelConfig(BaseConfig):


ServerType = Literal["vllm", "openai"]
ClientType = Literal["openai_chat_completions", "openai_chat_completions_token"]


class ElasticConfig(BaseConfig):
Expand Down Expand Up @@ -66,6 +67,17 @@ class ClientConfig(BaseConfig):
),
] = 1200

client_type: Annotated[
ClientType,
Field(
description=(
"Verifiers client type used by the orchestrator. "
"Use openai_chat_completions for standard /v1/chat/completions calls, "
"or openai_chat_completions_token for vLLM /chat/completions/tokens routing on multi-turn generation."
),
),
] = "openai_chat_completions_token"

base_url: Annotated[
list[str],
Field(
Expand All @@ -87,6 +99,24 @@ class ClientConfig(BaseConfig):
),
] = {}

sampling_overrides: Annotated[
dict[str, Any],
Field(
description=(
'Top-level request fields to hardcode/override on generation requests (e.g. {"logprobs": true}).'
),
),
] = {}

extra_body_overrides: Annotated[
dict[str, Any],
Field(
description=(
'extra_body fields to hardcode/override on generation requests (e.g. {"return_token_ids": true}).'
),
),
] = {}

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 CHANGELOG entry for new config fields

Low Severity

This PR adds three new config fields to ClientConfig in src/prime_rl/utils/config.pyclient_type, sampling_overrides, and extra_body_overrides — but the diff does not include a corresponding CHANGELOG.md update. Per project rules, any PR that modifies configuration structures (including added fields) in src/prime_rl/utils/config.py must update the changelog.

Fix in Cursor Fix in Web

Triggered by project rule: BugBot Instructions


skip_model_check: Annotated[
bool,
Field(
Expand Down
6 changes: 6 additions & 0 deletions src/prime_rl/utils/elastic.py
Original file line number Diff line number Diff line change
Expand Up @@ -165,9 +165,12 @@ def clients(self) -> list[vf.ClientConfig]:
setup_clients(
ClientConfig(
timeout=self.client_config.timeout,
client_type=self.client_config.client_type,
base_url=urls,
api_key_var=self.client_config.api_key_var,
headers=self.client_config.headers,
sampling_overrides=self.client_config.sampling_overrides,
extra_body_overrides=self.client_config.extra_body_overrides,
)
)
if urls
Expand Down Expand Up @@ -203,9 +206,12 @@ async def _create_admin_client(self, ip: str) -> AsyncClient:
url = self._build_url(ip)
config = ClientConfig(
timeout=self.client_config.timeout,
client_type=self.client_config.client_type,
base_url=[f"{url}/v1"],
api_key_var=self.client_config.api_key_var,
headers=self.client_config.headers,
sampling_overrides=self.client_config.sampling_overrides,
extra_body_overrides=self.client_config.extra_body_overrides,
)
return setup_admin_clients(config)[0]

Expand Down
55 changes: 55 additions & 0 deletions tests/unit/orchestrator/test_sampling_args.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
from prime_rl.orchestrator.config import EvalSamplingConfig, SamplingConfig
from prime_rl.orchestrator.eval_utils import get_eval_sampling_args
from prime_rl.orchestrator.utils import get_sampling_args
from prime_rl.utils.config import ClientConfig


def test_get_sampling_args_enforces_token_outputs_for_non_token_client():
sampling_config = SamplingConfig(
min_tokens=4,
repetition_penalty=1.1,
extra_body={"custom_flag": True},
)
client_config = ClientConfig(client_type="openai_chat_completions")

sampling_args = get_sampling_args(sampling_config, temperature=0.7, client_config=client_config)

assert sampling_args["logprobs"] is True
assert sampling_args["extra_body"]["return_token_ids"] is True
assert sampling_args["extra_body"]["custom_flag"] is True
assert sampling_args["temperature"] == 0.7


def test_get_sampling_args_applies_client_overrides():
sampling_config = SamplingConfig(extra_body={"top_k": 32, "custom": "from-sampling"})
client_config = ClientConfig(
sampling_overrides={"seed": 123, "logprobs": False},
extra_body_overrides={"top_k": 8, "return_token_ids": False, "trace": "enabled"},
)

sampling_args = get_sampling_args(sampling_config, temperature=1.0, client_config=client_config)

assert sampling_args["seed"] == 123
assert sampling_args["extra_body"]["top_k"] == 8
assert sampling_args["extra_body"]["trace"] == "enabled"
assert sampling_args["extra_body"]["custom"] == "from-sampling"
assert sampling_args["logprobs"] is True
assert sampling_args["extra_body"]["return_token_ids"] is True


def test_get_eval_sampling_args_applies_client_overrides():
eval_sampling_config = EvalSamplingConfig(
top_k=4,
extra_body={"trace": "from-eval"},
)
client_config = ClientConfig(
sampling_overrides={"seed": 99},
extra_body_overrides={"trace": "from-client", "return_token_ids": True},
)

sampling_args = get_eval_sampling_args(eval_sampling_config, client_config)

assert sampling_args["seed"] == 99
assert sampling_args["extra_body"]["top_k"] == 4
assert sampling_args["extra_body"]["trace"] == "from-client"
assert sampling_args["extra_body"]["return_token_ids"] is True
12 changes: 11 additions & 1 deletion tests/unit/utils/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,8 @@
import httpx
import pytest

from prime_rl.utils.client import _is_retryable_lora_error, load_lora_adapter
from prime_rl.utils.client import _is_retryable_lora_error, load_lora_adapter, setup_clients
from prime_rl.utils.config import ClientConfig


def test_is_retryable_lora_error_returns_true_for_404():
Expand Down Expand Up @@ -83,3 +84,12 @@ def test_load_lora_adapter_raises_non_retryable_error_immediately():

assert exc_info.value.response.status_code == 400
assert mock_client.post.call_count == 1


def test_setup_clients_uses_configured_client_type():
config = ClientConfig(base_url=["http://localhost:8000/v1"], client_type="openai_chat_completions")

clients = setup_clients(config)

assert len(clients) == 1
assert clients[0].client_type == "openai_chat_completions"
Loading