Skip to content
Merged
Show file tree
Hide file tree
Changes from 5 commits
Commits
Show all changes
16 commits
Select commit Hold shift + click to select a range
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 deps/verifiers
Submodule verifiers updated 74 files
+1 −1 .pre-commit-config.yaml
+91 −51 assets/lab/environments/AGENTS.md
+3 −2 docs/mint.json
+2 −2 docs/v1/architecture.md
+1 −1 docs/v1/evaluation.md
+44 −0 docs/v1/gepa.md
+1 −3 docs/v1/getting_started.md
+3 −3 docs/v1/harbor.md
+7 −11 docs/v1/overview.md
+61 −55 docs/v1/tasksets.md
+91 −51 environments/AGENTS.md
+1 −1 environments/alphabet_sort_v1/alphabet_sort_v1/servers/user.py
+10 −13 environments/color_codeword_v1/color_codeword_v1/servers/user.py
+5 −11 environments/color_codeword_v1/color_codeword_v1/taskset.py
+13 −22 environments/mmmu_v1/mmmu_v1/taskset.py
+3 −3 scripts/sync.py
+4 −4 skills/brainstorm/SKILL.md
+12 −12 skills/browse-environments/SKILL.md
+9 −9 skills/create-environments/SKILL.md
+6 −6 skills/evaluate-environments/SKILL.md
+5 −5 skills/evaluate-environments/references/REFERENCE.md
+5 −5 skills/train-with-environments/SKILL.md
+1 −1 tests/v1/fixtures/echo_user_sim_v1.py
+7 −7 tests/v1/test_color_codeword_tasks.py
+9 −1 verifiers/v1/__init__.py
+2 −4 verifiers/v1/cli/dashboard/eval.py
+9 −19 verifiers/v1/cli/dashboard/replay.py
+15 −8 verifiers/v1/cli/dashboard/validate.py
+3 −17 verifiers/v1/cli/debug.py
+0 −1 verifiers/v1/cli/eval/resume.py
+25 −18 verifiers/v1/cli/eval/runner.py
+1 −1 verifiers/v1/cli/gepa.py
+4 −5 verifiers/v1/cli/replay.py
+13 −0 verifiers/v1/decorators.py
+52 −25 verifiers/v1/env.py
+1 −1 verifiers/v1/errors.py
+5 −8 verifiers/v1/gepa/config.py
+2 −4 verifiers/v1/gepa/runner.py
+4 −1 verifiers/v1/graph.py
+3 −0 verifiers/v1/harnesses/__init__.py
+2 −1 verifiers/v1/harnesses/default/program.py
+16 −6 verifiers/v1/harnesses/null/program.py
+3 −0 verifiers/v1/harnesses/pi/__init__.py
+264 −0 verifiers/v1/harnesses/pi/harness.py
+86 −8 verifiers/v1/interception/__init__.py
+68 −0 verifiers/v1/interception/base.py
+122 −73 verifiers/v1/interception/pool.py
+78 −108 verifiers/v1/interception/server.py
+40 −0 verifiers/v1/interception/tunnel/__init__.py
+50 −0 verifiers/v1/interception/tunnel/base.py
+42 −0 verifiers/v1/interception/tunnel/custom.py
+56 −0 verifiers/v1/interception/tunnel/prime.py
+16 −6 verifiers/v1/legacy.py
+1 −1 verifiers/v1/loaders.py
+64 −42 verifiers/v1/mcp/launch.py
+9 −2 verifiers/v1/mcp/server.py
+38 −37 verifiers/v1/rollout.py
+1 −8 verifiers/v1/runtimes/__init__.py
+17 −118 verifiers/v1/runtimes/base.py
+12 −1 verifiers/v1/runtimes/docker.py
+0 −7 verifiers/v1/runtimes/limiters.py
+3 −3 verifiers/v1/runtimes/modal.py
+8 −10 verifiers/v1/runtimes/prime.py
+27 −7 verifiers/v1/serve/client.py
+1 −1 verifiers/v1/serve/pool.py
+23 −45 verifiers/v1/serve/server.py
+24 −6 verifiers/v1/serve/types.py
+100 −0 verifiers/v1/session.py
+11 −0 verifiers/v1/task.py
+9 −1 verifiers/v1/taskset.py
+29 −11 verifiers/v1/tasksets/harbor/taskset.py
+19 −25 verifiers/v1/tasksets/lean/taskset.py
+4 −2 verifiers/v1/tasksets/textarena/taskset.py
+9 −0 verifiers/v1/utils/image.py
11 changes: 9 additions & 2 deletions src/prime_rl/orchestrator/dispatcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -386,6 +386,7 @@ def next_fresh_group(self, kind: RolloutKind, envs) -> GroupState | None:
kind=kind,
env_name=env_name,
task_idx=example["task_idx"],
task=example.get("task"),
rollouts_to_schedule=group_size,
target_rollouts=group_size,
eval_step=eval_step,
Expand Down Expand Up @@ -436,17 +437,23 @@ async def schedule_group_rollout(self, group_id: uuid.UUID, group: GroupState) -
else:
cache_salt = None

# A v1 env takes the task itself; the legacy bridge is addressed by dataset row.
if group.task is not None:
addressing = {"task_data": group.task.data.model_dump(mode="json")}
else:
addressing = {"task_idx": group.task_idx}

if env.requires_group_scoring:
permits = group.rollouts_to_schedule
group.rollouts_to_schedule = 0
await self.acquire(permits)
task: asyncio.Task = asyncio.create_task(
env.run_group(
client=client,
task_idx=group.task_idx,
model_name=model_name,
group_size=permits,
cache_salt=cache_salt,
**addressing,
)
)
else:
Expand All @@ -456,9 +463,9 @@ async def schedule_group_rollout(self, group_id: uuid.UUID, group: GroupState) -
task = asyncio.create_task(
env.run_rollout(
client=client,
task_idx=group.task_idx,
model_name=model_name,
cache_salt=cache_salt,
**addressing,
)
)

Expand Down
103 changes: 73 additions & 30 deletions src/prime_rl/orchestrator/envs.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,15 +2,19 @@

Each ``Env`` owns a v1 ``EnvServer`` (spawned as a child process, or an
external one given by ``config.address``) and an ``EnvClient`` to drive it. The
orchestrator never *runs* an environment: it asks the server for ``info``
(``num_tasks`` + whether group scoring is needed), then runs rollouts purely by
**task index**. The server returns a ``Trace`` (a plain ``model_dump`` — derived values are
orchestrator never *runs* an environment — the harness and runtime live only in the
server — but it does own the *taskset*: a v1 env's tasks are loaded here, once, and each
dispatched rollout ships its task's data on the request (``task_data``); the server
pydantic-validates it into the taskset's declared ``TaskData`` type and runs it. That
keeps the server (and every worker in its pool) stateless about data — no per-worker
dataset loads, no idx-addressed task cache — and gives the orchestrator real tasks to
cycle, shuffle, and filter. Only the legacy (v0) bridge, whose dataset genuinely lives
server-side, is still driven by ``task_idx`` (its count comes from ``info``).

The server returns a ``Trace`` (a plain ``model_dump`` — derived values are
properties, not serialized) which we validate into a ``Trace[WireTaskData]`` — a real ``vf.Trace``
(never a loose dict) whose task keeps the env's
task-specific fields as extras (``WireTaskData`` allows them). The orchestrator never imports the
env package: the env's *type* and *runtime* both live only in the server, and the orchestrator
drives it purely by task index. (Nothing here reads typed env task fields — only ``task.idx``
and a full ``task.model_dump``, both of which ``WireTaskData`` preserves.)
(never a loose dict) whose task keeps the env's task-specific fields as extras
(``WireTaskData`` allows them).
"""

from __future__ import annotations
Expand All @@ -22,6 +26,7 @@
import queue
import sys
from collections.abc import Iterator, Sequence
from itertools import islice
from multiprocessing.process import BaseProcess
from pathlib import Path
from typing import Generic, TypeVar
Expand All @@ -40,9 +45,8 @@
# task.idx + task.model_dump).
ROLLOUT_TYPE = Rollout[vf.WireTaskData]

# Max wait for a spawned env server to bind and report its address. The child
# loads the taskset (possibly downloading a dataset) before reporting, so this
# is generous.
# Max wait for a spawned env server to bind and report its address. A legacy
# child loads its dataset before reporting, so this is generous.
ENV_SERVER_SPAWN_TIMEOUT = 600.0


Expand Down Expand Up @@ -78,14 +82,20 @@ def _run_env_server(


class Env:
"""Wraps a v1 env server + client. The orchestrator never loads the env."""
"""Wraps a v1 env server + client. The orchestrator owns the taskset (loaded once,
client-side); the server owns harness execution."""

def __init__(self, config: EnvConfig):
self.config = config
self.sampling_args: dict = {}
self.num_tasks: int | None = 0
"""Task count reported by the server; ``None`` means the taskset is infinite."""
"""Task count; ``None`` means the taskset is infinite."""
self.requires_group_scoring: bool = False
self.tasks: list[vf.Task] | None = None
"""A finite v1 taskset's tasks, materialized at ``start()``. None for legacy
(dataset lives on the server) and for infinite tasksets (see ``task_iter``)."""
self.task_iter: Iterator[vf.Task] | None = None
"""An infinite v1 taskset's generator — pulled per example, never materialized."""
self._env_client: EnvClient | None = None
self._env_server_process: BaseProcess | None = None

Expand All @@ -100,24 +110,35 @@ def env_client(self) -> EnvClient:
return self._env_client

async def start(self, log_dir: Path, log_level: str | None = None, json_logging: bool = False) -> None:
"""Spawn the env server (if needed), connect, and cache its ``info``."""
"""Spawn the env server (if needed), connect, and load the taskset client-side
(legacy instead asks the server for ``info`` — its dataset is server-side)."""
external = self.config.address is not None
address = self.config.address or await self._spawn(log_dir, log_level or "INFO", json_logging)
get_logger().debug(f"Connecting {self.name} to env server {address}")
self._env_client = EnvClient(address=address)
# A spawned server already reported its address *after* binding + loading,
# so it's up — the untimed ``info`` below is enough. An external server has
# no such handshake, so poll until it answers before we block on ``info``.
# A spawned server already reported its address *after* binding, so it's up. An
# external server has no such handshake, so poll until it answers.
if external:
await self.env_client.wait_for_server_startup()
info = await self.env_client.info()
self.num_tasks = info.num_tasks
self.requires_group_scoring = info.requires_group_scoring
if self.config.is_legacy:
info = await self.env_client.info()
self.num_tasks = info.num_tasks
self.requires_group_scoring = info.requires_group_scoring
else:
taskset = vf.load_taskset(self.config.taskset)
self.requires_group_scoring = vf.has_decorated(type(taskset).task_type(), "group_reward")
if type(taskset).INFINITE:
self.task_iter = iter(taskset.load())
self.num_tasks = None
else:
# Materialize off the event loop — load() may pull a dataset.
self.tasks = await asyncio.to_thread(lambda: list(taskset.load()))
self.num_tasks = len(self.tasks)
num_tasks = self.num_tasks if self.num_tasks is not None else "infinite"
get_logger().info(f"Env {self.name} ready: num_tasks={num_tasks} group_scoring={self.requires_group_scoring}")

async def _spawn(self, log_dir: Path, log_level: str, json_logging: bool) -> str:
"""Spawn a v1 EnvServer child process (it loads the env; we never do).
"""Spawn a v1 EnvServer child process (it runs harnesses; the tasks come from us).
The server binds an OS-assigned port (``:0``) and reports the concrete
address back over a queue — no free-port guess, no TOCTOU race. Its output
goes to ``<log_dir>/<name>.log`` (``log_dir`` is already the train/eval-split
Expand Down Expand Up @@ -169,10 +190,17 @@ def _sampling(self, cache_salt: str | None) -> vf.SamplingConfig:
return vf.SamplingConfig(**sampling)

async def run_rollout(
self, client: vf.ClientConfig, task_idx: int, model_name: str, cache_salt: str | None
self,
client: vf.ClientConfig,
model_name: str,
cache_salt: str | None,
task_data: dict | None = None,
task_idx: int | None = None,
) -> Rollout:
"""Run a single rollout for ``task_idx``; return a typed Trace."""
"""Run a single rollout; return a typed Trace. A v1 env takes the task itself
(``task_data``); the legacy bridge is addressed by dataset row (``task_idx``)."""
wire = await self.env_client.run_rollout(
task_data=task_data,
task_idx=task_idx,
client=client,
model=model_name,
Expand All @@ -181,10 +209,18 @@ async def run_rollout(
return ROLLOUT_TYPE.model_construct(**dict(wire))

async def run_group(
self, client: vf.ClientConfig, task_idx: int, model_name: str, group_size: int, cache_salt: str | None
self,
client: vf.ClientConfig,
model_name: str,
group_size: int,
cache_salt: str | None,
task_data: dict | None = None,
task_idx: int | None = None,
) -> list[Rollout]:
"""Run a group of rollouts for ``task_idx`` (group-scoring envs); return typed Traces."""
"""Run a group of rollouts of one task (group-scoring envs); return typed Traces.
Task addressing as in :meth:`run_rollout`."""
wires = await self.env_client.run_group(
task_data=task_data,
task_idx=task_idx,
n=group_size,
client=client,
Expand Down Expand Up @@ -220,13 +256,20 @@ def __init__(self, config: EvalEnvConfig):

async def start(self, log_dir: Path, log_level: str | None = None, json_logging: bool = False) -> None:
await super().start(log_dir=log_dir, log_level=log_level, json_logging=json_logging)
n = self.config.num_examples
if self.num_tasks is None:
if self.config.num_examples < 0:
if n < 0:
raise ValueError(f"Eval env {self.name} has an infinite taskset — set num_examples to bound it")
n = self.config.num_examples
else:
n = self.num_tasks if self.config.num_examples < 0 else min(self.config.num_examples, self.num_tasks)
self.examples = [{"task_idx": i} for i in range(n)]
assert self.task_iter is not None
# A fixed eval set off the generator, pulled once and reused every epoch.
tasks = list(islice(self.task_iter, n))
elif self.tasks is not None:
tasks = self.tasks if n < 0 else self.tasks[:n]
else: # legacy: the dataset lives on the server — address it by row
count = self.num_tasks if n < 0 else min(n, self.num_tasks)
self.examples = [{"task_idx": i} for i in range(count)]
return
self.examples = [{"task_idx": task.data.idx, "task": task} for task in tasks]


EnvT = TypeVar("EnvT", bound=Env)
Expand Down
40 changes: 26 additions & 14 deletions src/prime_rl/orchestrator/train_source.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,18 @@
"""TrainSource: weighted round-robin across train envs, infinite pull.

Weights are each env's configured ``ratio`` (default 1, i.e. equal weight
per env). A finite env serves a shuffled task-index table, reshuffled on
cursor exhaustion; an infinite env (``num_tasks is None``) streams a
monotonic ``task_idx`` — the server generates tasks on demand, so every
pull is a fresh task and there are no epochs to shuffle."""
per env). A v1 env serves the tasks the orchestrator loaded client-side: a
finite one as a shuffled table (reshuffled on cursor exhaustion), an
infinite one (``num_tasks is None``) straight off its generator — every
pull is a fresh task and there are no epochs to shuffle. A legacy env's
dataset lives on its server, so it serves shuffled task *indices*."""

from __future__ import annotations

import random
from collections.abc import Iterator

import verifiers.v1 as vf

from prime_rl.orchestrator.envs import TrainEnvs

Expand All @@ -17,28 +21,36 @@ class TrainSource:
"""``next_example(available_permits)`` picks a weighted-RR env and
returns its next example (or ``None`` when the env's per-call permit
cost doesn't fit — the dispatch loop retries when permits free up).
Returned dicts carry ``env_name`` + ``task_idx``."""
Returned dicts carry ``env_name`` + ``task_idx`` (+ ``task`` for v1 envs,
whose data is shipped to the env server at dispatch)."""

def __init__(self, train_envs: TrainEnvs, *, seed: int | None) -> None:
self.rng = random.Random(seed)
self.envs = list(train_envs)
if not self.envs:
raise ValueError("TrainSource needs at least one train env")

# A finite env's shuffled index table; ``None`` for an infinite env,
# whose cursor alone is the (monotonic) task_idx stream.
# A finite env's shuffled example table; ``None`` for an infinite env,
# whose generator (``self.iters``) is pulled per example.
self.examples: dict[str, list[dict] | None] = {}
self.iters: dict[str, Iterator[vf.Task]] = {}
self.cursors: dict[str, int] = {}
# Group-scoring envs reserve ``group_size`` permits up front;
# per-rollout envs need 1
self.env_costs: dict[str, int] = {}
for env in self.envs:
# The orchestrator never loads the env: sample over the task-index
# range the server reported via info() (num_tasks; None = infinite).
if env.num_tasks is None:
assert env.task_iter is not None
self.examples[env.name] = None
else:
rows: list[dict] = [{"task_idx": i, "env_name": env.name} for i in range(env.num_tasks)]
self.iters[env.name] = env.task_iter
elif env.tasks is not None:
rows: list[dict] = [
{"task_idx": task.data.idx, "task": task, "env_name": env.name} for task in env.tasks
]
self.rng.shuffle(rows)
self.examples[env.name] = rows
Comment thread
mikasenghaas marked this conversation as resolved.
else: # legacy: sample over the index range the server reported via info()
rows = [{"task_idx": i, "env_name": env.name} for i in range(env.num_tasks)]
self.rng.shuffle(rows)
self.examples[env.name] = rows
self.cursors[env.name] = 0
Expand All @@ -53,9 +65,9 @@ def next_example(self, available_permits: int) -> dict | None:
return None
rows = self.examples[env_name]
cursor = self.cursors[env_name]
if rows is None: # infinite env: the cursor is the task_idx
self.cursors[env_name] = cursor + 1
return {"task_idx": cursor, "env_name": env_name}
if rows is None: # infinite env: pull the next generated task
task = next(self.iters[env_name])
return {"task_idx": task.data.idx, "task": task, "env_name": env_name}
if cursor >= len(rows):
self.rng.shuffle(rows)
cursor = 0
Expand Down
3 changes: 3 additions & 0 deletions src/prime_rl/orchestrator/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,9 @@ class GroupState:
task_idx: int
rollouts_to_schedule: int
target_rollouts: int
task: vf.Task | None = None
"""The group's task (v1 envs — its data is shipped on every dispatch). ``None`` for
legacy envs, which are addressed by ``task_idx`` alone."""
emitted: int = 0
eval_step: int | None = None
pinned_client: vf.ClientConfig | None = None
Expand Down
Loading