feat: add Verifiers rollout integration (#1739)

Co-authored-by: Tao Lin <tao.lin@radixark.ai>
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
will brown
2026-08-06 15:44:05 -07:00
committed by GitHub
co-authored by Tao Lin Claude Fable 5
parent 03491d0431
commit d2010d2929
10 changed files with 2269 additions and 1 deletions
+2 -1
View File
@@ -182,7 +182,8 @@
"pages": [
"user-guide/harbor",
"user-guide/openenv",
"user-guide/nemo-gym"
"user-guide/nemo-gym",
"user-guide/verifiers"
]
}
]
+1
View File
@@ -23,6 +23,7 @@ where the environment itself comes from:
| [NeMo-Gym](https://github.com/NVIDIA-NeMo/Gym) | agent function | [guide](/user-guide/nemo-gym) |
| [Strands Agents](https://strandsagents.com/) | generate function | [example](https://github.com/radixark/miles/tree/main/examples/experimental/strands_sglang) |
| [τ-bench](https://github.com/sierra-research/tau-bench) | generate function | [example](https://github.com/radixark/miles/tree/main/examples/experimental/tau-bench) |
| [Verifiers (Prime Intellect)](https://github.com/PrimeIntellect-ai/verifiers) | rollout function | [guide](/user-guide/verifiers) |
Sandbox providers are a different axis: they provision the task containers
*inside* a connector rather than occupying a rollout layer.
+122
View File
@@ -0,0 +1,122 @@
---
title: Verifiers (Prime Intellect)
description: Train on Prime Intellect Verifiers environments with Miles.
---
Miles can train on a Verifiers environment in place of a prompt dataset. The
integration requires Python 3.11 or newer and Verifiers 0.2.0. Verifiers 0.2.1
requires OpenAI 2.9 or newer, while SGLang 0.5.15 pins OpenAI 2.6.1.
## Install
Install the adapter's dependencies and the Prime CLI:
```bash
pip install -r examples/experimental/verifiers/requirements.txt
uv tool install prime
```
The recommended workspace keeps local environment packages under `./environments`:
```text
workspace/
environments/
my-environment/
```
From the workspace root, install a local environment by name. For an environment from
the Environments Hub, authenticate and use its `user/environment` ID:
```bash
# Local: ./environments/my-environment
prime env install my-environment
# Environments Hub
prime login
prime env install user/my-environment
```
## Configure
Create a Verifiers `EnvConfig` TOML file. A minimal config selects a taskset:
```toml
[taskset]
id = "gsm8k-v1"
```
The config may also define the harness, runtime, judges, retries, and environment
limits supported by Verifiers. Verifiers applies per-rollout and group rewards before
the completed traces are returned to Miles.
The integration implements Verifiers' V1 environment contract. Legacy V0 environment
configs are rejected during startup.
## Run
The integration is a rollout function under
[`examples/experimental/verifiers`](https://github.com/radixark/miles/tree/main/examples/experimental/verifiers);
its launcher wires everything up:
```bash
python examples/experimental/verifiers/run.py --verifiers-config /path/to/verifiers.toml
```
The launcher selects the adapter with `--rollout-function-path`, turns off Miles
prompt-data loading with `--disable-rollout-global-dataset`, and points
`VERIFIERS_CONFIG` at the file. This uses the configured taskset instead of Miles
prompt data. Environment behavior comes from the Verifiers config, while Miles
continues to own the model, sampling, batching, concurrency, reward hooks, and
optimizer settings. The Renderers library formats environment messages with Miles'
model and tokenizer settings.
The standard Miles rollout options keep their existing meaning:
| Miles option | Verifiers behavior |
|---|---|
| `--rollout-batch-size` | Number of task groups returned by each training rollout |
| `--n-samples-per-prompt` | Rollouts per training task |
| `--n-samples-per-eval-prompt` | Rollouts per evaluation task |
| `--rollout-shuffle` / `--rollout-seed` | Finite taskset order and sampling seeds |
| `--rollout-*` / `--eval-*` sampling options | Sampling and context limits |
| `--apply-chat-template-kwargs` | Typed template options passed to renderers |
| `--sglang-server-concurrency` | Physical engine capacity |
| Miles reward and filtering options | Applied after Verifiers scoring using the standard Miles hooks |
Evaluation covers every task in the taskset. Training cycles the taskset and advances
from the current Miles rollout ID when a run resumes.
## Environment Support
The adapter supports V1 environments that use the Chat Completions dialect with
text-only Renderers inputs. Tools require a model-specific renderer; use a registered
model identity in `--hf-checkpoint` or the existing `--sglang-tokenizer-path` option.
User simulators, multi-turn episodes, environment runtimes, per-rollout rewards, and
group rewards run through the standard Verifiers environment lifecycle.
Verifiers group rewards apply during both training and evaluation. Miles
`--group-rm` hooks remain training-only, matching the standard Miles rollout path.
## Limitations
`--eval-interval` works through the launcher, which evaluates the whole taskset at
that interval. Miles asserts that eval datasets are configured whenever the flag is
set, so the launcher passes a placeholder `--eval-prompt-data` naming the taskset and
pointing at its EnvConfig; the adapter serves evaluation, so the built-in loader never
opens that path.
`--partial-rollout` is not supported. A Verifiers episode owns live harness and
environment state and has no contract for resuming a partially executed episode. The
adapter rejects this combination when it is constructed, before any episode runs.
`--chat-template-path` is also rejected because Renderers owns message formatting for
Verifiers environments. Use the checkpoint's native template and
`--apply-chat-template-kwargs` instead.
Streaming model requests, Responses and Anthropic dialects, multimodal inputs, OPD,
routing replay, and indexer replay are not supported by the transport. The adapter
rejects the corresponding CLI options at startup.
Traces with multiple graph branches, including compaction, are rejected. Miles does
not currently preserve a trace's rollout-group boundary when it flattens multiple
training samples, which would make group-relative advantages incorrect.
+105
View File
@@ -0,0 +1,105 @@
# Verifiers (Prime Intellect) rollout integration
Train on a [Verifiers](https://github.com/PrimeIntellect-ai/verifiers) V1
environment instead of a Miles prompt dataset. Verifiers owns grouped episode
execution and reward computation; Miles keeps the model, sampling, engines and
weight updates, filtering, advantages, and the optimizer.
The adapter is a **rollout function** (`verifiers_rollout.py`): it replaces
Miles' batch-orchestration layer, runs `n` rollouts per task through Verifiers,
and returns ordinary Miles `Sample` groups. Renderers renders messages to token
ids and `MilesSGLangTransport` translates its wire format to Miles' SGLang
`/generate`, so training sees the exact sampled token ids and logprobs.
Requires Python 3.11+ and Verifiers 0.2.0. (Verifiers 0.2.1 requires OpenAI
2.9, while SGLang pins OpenAI 2.6.1.)
## Install
```bash
pip install -r examples/experimental/verifiers/requirements.txt
uv tool install prime
```
Install the environment itself with the Prime CLI, from a workspace that keeps
local environment packages under `./environments`:
```bash
# Local: ./environments/my-environment
prime env install my-environment
# Environments Hub
prime login
prime env install user/my-environment
```
## Configure
Write a Verifiers `EnvConfig` TOML. A minimal config selects a taskset:
```toml
[taskset]
id = "gsm8k-v1"
```
It may also define the harness, runtime, judges, retries, and environment
limits Verifiers supports. Legacy V0 configs are rejected at startup.
## Run
```bash
python examples/experimental/verifiers/run.py --verifiers-config /path/to/verifiers.toml
```
The launcher points `VERIFIERS_CONFIG` at the file and selects the adapter with
`--rollout-function-path verifiers_rollout.VerifiersRolloutFn` (the
`generate_rollout` function entry without
`MILES_EXPERIMENTAL_ROLLOUT_REFACTOR=1`) plus
`--disable-rollout-global-dataset`. To wire a hand-rolled command, set the same
three things and put this directory on `PYTHONPATH`.
Standard Miles options keep their meaning:
| Miles option | Verifiers behavior |
|---|---|
| `--rollout-batch-size` | Task groups per training rollout |
| `--n-samples-per-prompt` | Rollouts per training task |
| `--n-samples-per-eval-prompt` | Rollouts per evaluation task |
| `--rollout-shuffle` / `--rollout-seed` | Taskset order and sampling seeds |
| `--rollout-*` / `--eval-*` sampling options | Sampling and context limits |
| `--apply-chat-template-kwargs` | Typed template options passed to renderers |
| `--sglang-server-concurrency` | Physical engine capacity |
| Reward and filtering options | Applied after Verifiers scoring |
`--hf-checkpoint` and `--sglang-tokenizer-path` provide the model and renderer
identity. Tools need a renderer registered for that identity, so pass a
registered model id in one of them when the checkpoint is a local snapshot.
Evaluation covers every task in the taskset; training cycles it and advances
from the current rollout id when a run resumes.
## Unsupported
The adapter raises at construction rather than producing wrong data:
- `--partial-rollout` — a live episode has no resume contract.
- `--chat-template-path` — renderers owns message formatting; use
`--apply-chat-template-kwargs`.
- `--multimodal-keys`, `--use-opd`, routing replay, indexer replay — the
transport does not carry that token metadata.
- Streaming, Responses/Anthropic dialects, and auxiliary relay routes.
- Traces with multiple graph branches, including compaction: Miles does not
preserve a trace's rollout-group boundary when flattening, which would make
group-relative advantages wrong.
## Evaluation
`run.py --eval-interval N` evaluates the whole taskset every N training
rollouts. Miles asserts that eval datasets are configured whenever
`--eval-interval` is set, so the launcher works around that by naming the
taskset and pointing the placeholder `--eval-prompt-data` at the EnvConfig it is
defined in; the adapter serves evaluation, so the built-in loader never opens
that path. Scoping the assertion Miles-side is worth revisiting once a second
rollout function owns its evaluation set.
Group rewards rank rollouts within a task, so evaluation needs
`--n-samples-per-eval-prompt >= 2` (the launcher's default).
@@ -0,0 +1,4 @@
openai-agents<0.5
renderers>=0.1.8
# 0.2.1 requires OpenAI>=2.9, while SGLang pins OpenAI==2.6.1.
verifiers>=0.2.0,<0.2.1
+182
View File
@@ -0,0 +1,182 @@
"""Verifiers launcher (Qwen3-0.6B on code-golf-v1): Miles <-> Verifiers V1.
Defaults reproduce the two-GPU smoke configuration: 3 GRPO steps against the
`code-golf-v1` taskset. Scale --num-rollout and the batch sizes for real
training, and point --verifiers-config at your own EnvConfig TOML.
The Verifiers environment must already be installed in this interpreter, e.g.
`prime env install code-golf-v1` from a workspace with ./environments (see
README.md).
Usage:
python run.py
python run.py --verifiers-config /path/to/verifiers.toml --num-rollout 50
python run.py --eval-interval 5 # evaluate the whole taskset every 5 rollouts
"""
import os
from dataclasses import dataclass
from pathlib import Path
import typer
import miles.utils.external_utils.command_utils as U
SCRIPT_DIR = Path(__file__).resolve().parent
@dataclass
class ScriptArgs(U.ExecuteTrainConfig):
megatron_model_type: str = "qwen3-0.6B"
num_gpus_per_node: int = 2
megatron_path: str = "/root/Megatron-LM"
# Paths
skip_prepare: bool = False
model_name: str = "Qwen3-0.6B"
hf_checkpoint: str = "/root/models/Qwen3-0.6B"
# Renderers resolves its renderer from the model identity, which a local
# snapshot path does not carry; keep the registered id here.
sglang_tokenizer_path: str = "Qwen/Qwen3-0.6B"
ref_load: str = "/root/models/Qwen3-0.6B_torch_dist"
save_dir: str = "/root/Qwen3-0.6B_verifiers/"
# The Verifiers EnvConfig TOML. Written below when left at the default.
verifiers_config: str = "/root/verifiers-code-golf.toml"
taskset_id: str = "code-golf-v1"
# Training settings (smoke scale)
rollout_max_response_len: int = 512
rollout_max_context_len: int = 2048
num_rollout: int = 3
rollout_batch_size: int = 3
n_samples_per_prompt: int = 4
global_batch_size: int = 12
# Evaluation over the whole taskset, every N training rollouts. 0 disables it.
# Group rewards need at least two rollouts per task to rank.
eval_interval: int = 0
n_samples_per_eval_prompt: int = 2
def prepare(args: ScriptArgs):
U.exec_command(f"hf download Qwen/{args.model_name} --local-dir {args.hf_checkpoint}")
U.convert_checkpoint(
model_name=args.model_name,
megatron_model_type=args.megatron_model_type,
num_gpus_per_node=args.num_gpus_per_node,
dir_dst=str(Path(args.hf_checkpoint).parent),
hf_checkpoint=args.hf_checkpoint,
megatron_path=args.megatron_path,
)
def execute(args: ScriptArgs):
config_path = Path(args.verifiers_config)
if not config_path.exists():
config_path.write_text(f'[taskset]\nid = "{args.taskset_id}"\n')
ckpt_args = (
f"--hf-checkpoint {args.hf_checkpoint} "
f"--sglang-tokenizer-path {args.sglang_tokenizer_path} "
f"--ref-load {args.ref_load} "
f"--save {args.save_dir} "
"--save-interval 1000 "
)
# Verifiers owns the taskset, so Miles loads no prompt data; the rollout
# function plug-point selects the adapter, which resolves as a bare module
# because PYTHONPATH carries this directory into the rollout actor.
rollout_fn = (
"verifiers_rollout.VerifiersRolloutFn"
if os.environ.get("MILES_EXPERIMENTAL_ROLLOUT_REFACTOR") == "1"
else "verifiers_rollout.generate_rollout"
)
rollout_args = (
f"--rollout-function-path {rollout_fn} "
"--disable-rollout-global-dataset "
f"--num-rollout {args.num_rollout} "
f"--rollout-batch-size {args.rollout_batch_size} "
f"--n-samples-per-prompt {args.n_samples_per_prompt} "
f"--over-sampling-batch-size {args.rollout_batch_size} "
f"--rollout-max-response-len {args.rollout_max_response_len} "
f"--rollout-max-context-len {args.rollout_max_context_len} "
"--rollout-temperature 0.8 "
f"--rollout-num-gpus-per-engine 1 "
)
# Workaround: the taskset is the evaluation set, but Miles asserts that eval
# datasets are configured whenever --eval-interval is set, so name the taskset
# and point the placeholder at the config it is defined in -- the adapter serves
# eval, so the built-in loader never opens this path. Worth replacing with a
# Miles-side fix once a second rollout function owns its evaluation set.
eval_args = (
(
f"--eval-interval {args.eval_interval} "
f"--n-samples-per-eval-prompt {args.n_samples_per_eval_prompt} "
f"--eval-prompt-data verifiers-taskset {config_path} "
)
if args.eval_interval
else ""
)
grpo_args = (
"--advantage-estimator grpo "
f"--global-batch-size {args.global_batch_size} "
"--balance-data "
"--entropy-coef 0.0 "
"--eps-clip 0.2 "
"--eps-clip-high 0.28 "
)
optimizer_args = (
"--optimizer adam "
"--lr 1e-6 "
"--lr-decay-style constant "
"--weight-decay 0.1 "
"--adam-beta1 0.9 "
"--adam-beta2 0.98 "
)
perf_args = (
"--tensor-model-parallel-size 1 "
"--pipeline-model-parallel-size 1 "
"--context-parallel-size 1 "
"--use-dynamic-batch-size "
"--max-tokens-per-gpu 4096 "
"--no-gradient-accumulation-fusion "
)
sglang_args = "--sglang-mem-fraction-static 0.6 --sglang-enable-metrics "
misc_args = (
"--attention-backend flash "
"--colocate "
"--actor-num-nodes 1 "
f"--actor-num-gpus-per-node {args.num_gpus_per_node} "
)
U.execute_train(
train_args=(
f"{ckpt_args}{rollout_args}{eval_args}{grpo_args}{optimizer_args}{perf_args}{sglang_args}{misc_args}"
),
config=args,
num_gpus_per_node=args.num_gpus_per_node,
megatron_model_type=args.megatron_model_type,
megatron_path=args.megatron_path,
extra_env_vars={
"PYTHONPATH": f"{args.megatron_path}:{SCRIPT_DIR}:{U.repo_base_dir}",
"VERIFIERS_CONFIG": str(config_path),
},
)
@U.dataclass_cli
def main(args: ScriptArgs):
if not args.skip_prepare:
prepare(args)
execute(args)
if __name__ == "__main__":
typer.run(main)
@@ -0,0 +1,887 @@
"""Verifiers V1 rollout adapter: Verifiers owns grouped episode execution and
reward computation, Miles owns everything around it.
Wire it up with ``--rollout-function-path verifiers_rollout.VerifiersRolloutFn``
(or ``.generate_rollout`` without ``MILES_EXPERIMENTAL_ROLLOUT_REFACTOR``) plus
``--disable-rollout-global-dataset``, and point ``VERIFIERS_CONFIG`` at a
Verifiers ``EnvConfig`` TOML file. ``run.py`` in this directory does all of
that; see README.md.
"""
from __future__ import annotations
import asyncio
import logging
import os
import random
import sys
import uuid
from argparse import Namespace
from collections import OrderedDict
from collections.abc import Iterable
from importlib import metadata as importlib_metadata
from pathlib import Path
from types import SimpleNamespace
from typing import Any
import httpx
from packaging.version import InvalidVersion, Version
from miles.rollout.base_types import (
RolloutFnConstructorInput,
RolloutFnEvalInput,
RolloutFnEvalOutput,
RolloutFnInput,
RolloutFnOutput,
RolloutFnTrainInput,
RolloutFnTrainOutput,
)
from miles.rollout.filter_hub.base_types import MetricGatherer, call_dynamic_filter
from miles.rollout.generate_utils.prefill_logprobs import recompute_samples_rollout_logprobs_via_prefill
from miles.utils.lora import LORA_ADAPTER_NAME, is_lora_enabled
from miles.utils.types import Sample
logger = logging.getLogger(__name__)
_MIN_VERIFIERS_VERSION = Version("0.2.0")
_MAX_VERIFIERS_VERSION = Version("0.2.1")
_MIN_RENDERERS_VERSION = Version("0.1.8")
_UNSUPPORTED_ERROR_PREFIX = "Miles' Verifiers adapter does not support"
_CONFIG_ENV_VAR = "VERIFIERS_CONFIG"
def _load_config_data(path: str) -> dict[str, Any]:
config_path = Path(path)
if config_path.suffix.lower() != ".toml":
raise ValueError("--verifiers-config must point to a Verifiers TOML config.")
if sys.version_info < (3, 11):
raise _optional_dependency_error()
import tomllib
data = tomllib.loads(config_path.read_text())
if not isinstance(data, dict):
raise ValueError(f"{path} must contain a mapping at the root.")
return data
def _optional_dependency_error() -> RuntimeError:
return RuntimeError(
"Verifiers rollouts require Python 3.11+ and the optional dependencies. "
"Install Miles with `pip install -e '.[verifiers]'`."
)
def _installed_version(package: str) -> str:
try:
return importlib_metadata.version(package)
except importlib_metadata.PackageNotFoundError as error:
raise _optional_dependency_error() from error
def _check_version(package: str, raw_version: str, minimum: Version, maximum: Version | None = None) -> None:
try:
installed = Version(raw_version)
except InvalidVersion as error:
raise RuntimeError(f"Could not parse installed {package} version {raw_version!r}.") from error
if installed < minimum:
raise RuntimeError(f"Verifiers rollouts require {package}>={minimum}; found {installed}.")
if maximum is not None and installed >= maximum:
raise RuntimeError(
f"Verifiers rollouts require {package}>={minimum},<{maximum}; found {installed}. "
"Verifiers 0.2.1 requires OpenAI>=2.9, while SGLang 0.5.15 pins OpenAI==2.6.1."
)
def _import_verifiers():
if sys.version_info < (3, 11):
raise _optional_dependency_error()
_check_version(
"verifiers",
_installed_version("verifiers"),
_MIN_VERIFIERS_VERSION,
_MAX_VERIFIERS_VERSION,
)
_check_version("renderers", _installed_version("renderers"), _MIN_RENDERERS_VERSION)
try:
from verifiers.v1 import EnvConfig, Environment, ModelContext, SamplingConfig
from verifiers.v1.clients.train import TrainClient
from verifiers.v1.decorators import discover_decorated
from verifiers.v1.errors import OverlongPromptError, ProviderError
except ImportError as error:
raise _optional_dependency_error() from error
return SimpleNamespace(
EnvConfig=EnvConfig,
Environment=Environment,
discover_decorated=discover_decorated,
ModelContext=ModelContext,
OverlongPromptError=OverlongPromptError,
ProviderError=ProviderError,
SamplingConfig=SamplingConfig,
TrainClient=TrainClient,
)
def _renderer_identity(checkpoint: str) -> str | None:
from renderers.base import MODEL_RENDERER_MAP
if checkpoint in MODEL_RENDERER_MAP:
return checkpoint
path = Path(checkpoint)
candidates = []
for part in path.parts:
if part.startswith("models--"):
candidates.append(part.removeprefix("models--").replace("--", "/"))
candidates.extend(model_id for model_id in MODEL_RENDERER_MAP if model_id.rsplit("/", 1)[-1] == path.name)
matches = sorted(set(candidates) & MODEL_RENDERER_MAP.keys())
if not matches:
return None
renderer_names = {MODEL_RENDERER_MAP[model_id] for model_id in matches}
return matches[0] if len(renderer_names) == 1 else None
def _train_client(
runtime,
args: Namespace,
model: str,
pool_size: int,
*,
router_args: Namespace | None = None,
):
tokenizer_source = getattr(args, "sglang_tokenizer_path", None) or model
identity = _renderer_identity(model) or _renderer_identity(tokenizer_source)
# TrainClient uses one path for both tokenizer loading and renderer lookup.
# Miles checkpoints are commonly local snapshots, so keep local tokenizer
# files while restoring the canonical identity used by Renderers' registry.
class TrainClient(runtime.TrainClient):
@staticmethod
def _unsupported_request(kind: str):
return runtime.ProviderError(
f"{_UNSUPPORTED_ERROR_PREFIX} {kind}.",
status_code=400,
)
async def get_response(self, *args, **kwargs):
try:
return await super().get_response(*args, **kwargs)
except NotImplementedError as error:
raise runtime.ProviderError(
f"{_UNSUPPORTED_ERROR_PREFIX} this request: {error}",
status_code=400,
) from error
except ValueError as error:
if "does not support tools" not in str(error):
raise
raise runtime.ProviderError(
f"{_UNSUPPORTED_ERROR_PREFIX} tools with this renderer: {error} "
"Use a Renderers-registered model identity in "
"--hf-checkpoint or --sglang-tokenizer-path.",
status_code=400,
) from error
async def relay(self, *args, **kwargs):
raise self._unsupported_request("streaming requests")
async def relay_aux(self, *args, **kwargs):
raise self._unsupported_request("auxiliary dialect routes")
def _renderer_pool(self, requested_model, *, chat_template_kwargs=None):
if identity is None:
return super()._renderer_pool(
requested_model,
chat_template_kwargs=chat_template_kwargs,
)
if self._pool is None:
from renderers import RendererPool, create_renderer
from renderers.base import load_tokenizer
source = self.renderer_model_name or requested_model
def factory():
tokenizer = load_tokenizer(source)
tokenizer.name_or_path = identity
return create_renderer(
tokenizer,
self.config,
chat_template_kwargs=chat_template_kwargs,
)
self._pool = RendererPool(factory, size=self.pool_size)
return self._pool
return TrainClient(
MilesSGLangTransport(args, router_args=router_args),
pool_size=pool_size,
renderer_model_name=tokenizer_source,
)
def _generate_url(args: Namespace, endpoint: str = "/generate") -> str:
routers = getattr(args, "sglang_model_routers", None)
if routers and "default" in routers:
ip, port = routers["default"]
else:
ip, port = args.sglang_router_ip, args.sglang_router_port
return f"http://{ip}:{port}{endpoint}"
async def _sglang_worker_urls(args: Namespace) -> list[str]:
from miles.utils.http_utils import get
router_url = _generate_url(args).removesuffix("/generate")
if not getattr(args, "use_miles_router", False):
try:
response = await get(f"{router_url}/workers")
return [worker["url"] for worker in response["workers"]]
except Exception:
logger.debug("SGLang /workers lookup failed; trying Miles /list_workers.", exc_info=True)
response = await get(f"{router_url}/list_workers")
return list(response["urls"])
def _finish_reason(output: dict[str, Any]) -> str:
finish_reason = (output.get("meta_info") or {}).get("finish_reason")
if isinstance(finish_reason, dict):
finish_reason = finish_reason.get("type")
if finish_reason == "abort":
raise RuntimeError("SGLang aborted the Verifiers generation request.")
return finish_reason if finish_reason in {"stop", "length", "content_filter"} else "stop"
class MilesSGLangTransport:
"""Translate Renderers' /inference/v1/generate wire format to Miles' SGLang endpoint."""
def __init__(self, args: Namespace, *, router_args: Namespace | None = None):
self.args = args
self.router_args = router_args or args
self._seen_sessions: OrderedDict[str, None] = OrderedDict()
self._session_cache_size = 10_000
@property
def base_url(self) -> str:
# The refactored RolloutManager constructs rollout functions before it
# starts SGLang and fills in the router address.
return f"{_generate_url(self.router_args, '').rstrip('/')}/v1"
async def get(self, _path: str, **_kwargs) -> dict[str, list[Any]]:
# Renderers probes this endpoint for an engine context cap (max_model_len).
# Miles owns separate prompt and response limits, so the transport
# enforces them at POST time.
return {"data": []}
def _sampling_params(self, raw: dict[str, Any], prompt_len: int) -> dict[str, Any]:
values = dict(raw)
values.pop("logprobs", None)
for source, target in (
("max_tokens", "max_new_tokens"),
("min_tokens", "min_new_tokens"),
("seed", "sampling_seed"),
):
if source in values:
values[target] = values.pop(source)
if self.args.rollout_stop is not None:
values.setdefault("stop", self.args.rollout_stop)
if self.args.rollout_stop_token_ids is not None:
renderer_stops = list(values.get("stop_token_ids") or [])
values["stop_token_ids"] = list(dict.fromkeys([*renderer_stops, *self.args.rollout_stop_token_ids]))
values["skip_special_tokens"] = self.args.rollout_skip_special_tokens
values["no_stop_trim"] = True
values["spaces_between_special_tokens"] = False
values["n"] = 1
context_limit = self.args.rollout_max_context_len
response_limit = self.args.rollout_max_response_len
requested = int(values.get("max_new_tokens", response_limit))
if context_limit is not None:
requested = min(requested, context_limit - prompt_len)
values["max_new_tokens"] = min(requested, response_limit)
if values["max_new_tokens"] <= 0:
runtime = _import_verifiers()
raise runtime.OverlongPromptError(
f"prompt has {prompt_len} tokens, rollout_max_context_len={context_limit}"
)
return values
async def post(self, endpoint: str, *, body: dict[str, Any], options=None, **_kwargs) -> httpx.Response:
if body.get("features") is not None:
raise NotImplementedError("Miles Verifiers rollouts do not yet support multimodal renderer features.")
prompt_ids = list(body["token_ids"])
headers = dict((options or {}).get("headers") or {})
session_id = headers.get("X-Session-ID")
if session_id:
if session_id in self._seen_sessions:
self._seen_sessions.move_to_end(session_id)
else:
max_prompt_len = getattr(self.args, "rollout_max_prompt_len", None)
if max_prompt_len is not None and len(prompt_ids) > max_prompt_len:
runtime = _import_verifiers()
raise runtime.OverlongPromptError(
f"initial prompt has {len(prompt_ids)} tokens, rollout_max_prompt_len={max_prompt_len}"
)
self._seen_sessions[session_id] = None
if len(self._seen_sessions) > self._session_cache_size:
self._seen_sessions.popitem(last=False)
payload: dict[str, Any] = {
"input_ids": prompt_ids,
"sampling_params": self._sampling_params(body.get("sampling_params") or {}, len(prompt_ids)),
"return_logprob": True,
}
if is_lora_enabled(self.args):
payload["lora_path"] = LORA_ADAPTER_NAME
if body.get("priority") is not None:
payload["priority"] = body["priority"]
if body.get("cache_salt") is not None:
payload["extra_key"] = body["cache_salt"]
request_headers = None
if getattr(self.args, "sglang_router_policy", None) in ("consistent_hashing", "manual") and session_id:
request_headers = {"X-SMG-Routing-Key": session_id}
from miles.utils.http_utils import post
output = await post(_generate_url(self.router_args), payload, headers=request_headers)
meta_info = dict(output.get("meta_info") or {})
token_logprobs = list(meta_info.get("output_token_logprobs") or [])
completion_ids = [int(item[1]) for item in token_logprobs]
completion_logprobs = [float(item[0]) for item in token_logprobs]
expected = int(meta_info.get("completion_tokens", len(completion_ids)))
if len(completion_ids) != expected:
raise RuntimeError(
"SGLang generate response has mismatched completion token metadata: "
f"{len(completion_ids)} != {expected}"
)
response_body = {
"request_id": output.get("request_id") or f"vf-{uuid.uuid4().hex}",
"choices": [
{
"token_ids": completion_ids,
"logprobs": {"content": [{"logprob": value} for value in completion_logprobs]},
"finish_reason": _finish_reason(output),
}
],
}
request = httpx.Request("POST", endpoint)
return httpx.Response(200, json=response_body, request=request)
async def close(self) -> None:
return None
def _sample_status(trace) -> Sample.Status:
if trace.has_error:
return Sample.Status.FAILED
if trace.is_truncated:
return Sample.Status.TRUNCATED
return Sample.Status.COMPLETED
def _serialize_prompt(prompt):
if isinstance(prompt, list):
return [
message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message
for message in prompt
]
return prompt or ""
def _validate_group_reward_sample_counts(args: Namespace, tasks, discover_decorated) -> None:
if not any(discover_decorated(task, "group_reward") for task in tasks):
return
if getattr(args, "num_rollout", None) != 0 and args.n_samples_per_prompt < 2:
raise ValueError("Verifiers tasks with @group_reward require --n-samples-per-prompt >= 2.")
if getattr(args, "eval_interval", None) is not None and args.n_samples_per_eval_prompt < 2:
raise ValueError("Verifiers tasks with @group_reward require --n-samples-per-eval-prompt >= 2.")
def _raise_for_unsupported_trace_errors(traces) -> None:
for trace in traces:
if trace.error is not None and trace.error.message.startswith(_UNSUPPORTED_ERROR_PREFIX):
raise RuntimeError(trace.error.message)
def _branch_to_sample(args: Namespace, trace, branch, *, group_index: int, index: int) -> Sample:
tokens = list(branch.token_ids)
sampled_mask = list(branch.sampled_mask)
logprobs = list(branch.logprobs)
if len(tokens) != len(sampled_mask) or len(tokens) != len(logprobs):
raise ValueError(
f"Trace {trace.id} token metadata mismatch: "
f"tokens={len(tokens)}, mask={len(sampled_mask)}, logprobs={len(logprobs)}"
)
first_sampled = sampled_mask.index(True) if True in sampled_mask else len(tokens)
response_length = len(tokens) - first_sampled
reward = trace.reward if args.reward_key is None else {**trace.rewards, "reward": trace.reward}
task_data = trace.task.data
label = getattr(task_data, "label", None)
if label is None:
label = getattr(task_data, "answer", None)
metadata = {
"verifiers": {
"branch_index": branch.index,
"task_index": getattr(task_data, "idx", None),
"rewards": dict(trace.rewards),
"metrics": dict(trace.metrics),
"stop_condition": trace.stop_condition,
}
}
if trace.error is not None:
metadata["verifiers"]["error"] = trace.error.model_dump(mode="json", exclude_none=True)
sample = Sample(
group_index=group_index,
index=index,
prompt=_serialize_prompt(getattr(task_data, "prompt", "")),
tokens=tokens,
response=trace.last_reply,
response_length=response_length,
label=label,
reward=reward,
loss_mask=[int(value) for value in sampled_mask[first_sampled:]],
rollout_log_probs=logprobs[first_sampled:],
status=_sample_status(trace),
metadata=metadata,
routing_key=trace.id,
)
sample.validate()
return sample
def trace_to_samples(args: Namespace, trace, *, group_index: int, index_start: int) -> list[Sample]:
if not trace.branches:
error = trace.error.model_dump(mode="json", exclude_none=True) if trace.error is not None else None
logger.warning(
"Verifiers trace %s has no graph branches; omitting it from training (error=%s).",
trace.id,
error,
)
return []
if len(trace.branches) != 1:
raise NotImplementedError(
"Miles cannot yet preserve Verifiers trace groups when a rollout produces "
f"multiple graph branches (trace {trace.id} produced {len(trace.branches)})."
)
return [
_branch_to_sample(
args,
trace,
trace.branches[0],
group_index=group_index,
index=index_start,
)
]
def trace_to_sample(args: Namespace, trace, *, group_index: int, index: int) -> Sample:
samples = trace_to_samples(args, trace, group_index=group_index, index_start=index)
if len(samples) != 1:
raise ValueError(f"Verifiers trace {trace.id} produced {len(samples)} branches, expected one.")
return samples[0]
def _trace_metrics(traces) -> dict[str, float]:
if not traces:
return {}
return {
"verifiers/reward_mean": sum(trace.reward for trace in traces) / len(traces),
"verifiers/error_rate": sum(trace.has_error for trace in traces) / len(traces),
"verifiers/truncated_rate": sum(trace.is_truncated for trace in traces) / len(traces),
"verifiers/num_turns_mean": sum(trace.num_turns for trace in traces) / len(traces),
}
def _trace_eval_reward(trace, reward_key: str | None):
if reward_key is None:
return trace.reward
rewards = {**trace.rewards, "reward": trace.reward}
if trace.has_error:
return rewards.get(reward_key)
return rewards[reward_key]
def _flatten_samples(values: Iterable[Any]) -> list[Sample]:
flattened = []
for value in values:
if isinstance(value, list):
flattened.extend(_flatten_samples(value))
else:
flattened.append(value)
return flattened
def _make_eval_args(args: Namespace) -> Namespace:
eval_args = Namespace(**vars(args))
for eval_name, rollout_name in (
("eval_temperature", "rollout_temperature"),
("eval_top_p", "rollout_top_p"),
("eval_top_k", "rollout_top_k"),
("eval_max_response_len", "rollout_max_response_len"),
("eval_max_context_len", "rollout_max_context_len"),
):
if (value := getattr(args, eval_name, None)) is not None:
setattr(eval_args, rollout_name, value)
eval_args.rollout_max_prompt_len = args.eval_max_prompt_len
eval_args.rollout_min_new_tokens = args.eval_min_new_tokens
eval_args.reward_key = args.eval_reward_key or args.reward_key
return eval_args
def _config_path() -> str:
path = os.environ.get(_CONFIG_ENV_VAR)
if not path:
raise ValueError(
f"Verifiers rollouts need {_CONFIG_ENV_VAR} set to a Verifiers EnvConfig TOML file. "
"Launch through run.py in this directory, or forward the variable yourself "
"(Ray runtime_env for a hand-rolled command)."
)
return path
def _validate_args(args: Namespace) -> None:
"""Reject the Miles options this adapter cannot honor.
run.py never builds these combinations; this catches a hand-rolled command
before an episode runs and produces silently wrong training data.
"""
if getattr(args, "rollout_global_dataset", False):
raise ValueError(
"Verifiers rollouts replace Miles prompt data with the configured taskset; "
"pass --disable-rollout-global-dataset."
)
if args.partial_rollout:
raise ValueError(
"--partial-rollout is not supported for Verifiers because an episode "
"cannot be resumed from partially executed environment state."
)
if args.multimodal_keys is not None:
raise ValueError(
"--multimodal-keys is not supported by the Verifiers transport, which "
"handles text-only renderer inputs."
)
if args.chat_template_path is not None:
raise ValueError(
"--chat-template-path is not supported because renderers does not accept a custom "
"Jinja template. Use the checkpoint's template and --apply-chat-template-kwargs."
)
unsupported = [
flag
for enabled, flag in (
(args.use_opd, "--use-opd"),
(args.use_rollout_routing_replay, "--use-rollout-routing-replay"),
(getattr(args, "use_rollout_indexer_replay", False), "--use-rollout-indexer-replay"),
)
if enabled
]
if unsupported:
raise ValueError(
f"{', '.join(unsupported)} is not supported by the Verifiers SGLang transport, "
"which does not preserve its additional token metadata."
)
class VerifiersRolloutFn:
def __init__(self, input: RolloutFnConstructorInput):
runtime = _import_verifiers()
self.args = input.args
self.data_source = input.data_source
_validate_args(self.args)
self.config = runtime.EnvConfig.model_validate(_load_config_data(_config_path()))
if self.config.is_legacy:
raise ValueError("Miles' Verifiers integration supports V1 environment configs only.")
if self.config.harness.id == "codex":
raise ValueError(
"Miles' Verifiers adapter does not support the Codex harness because it uses the Responses dialect."
)
self.env = runtime.Environment(self.config)
self.model = self.args.hf_checkpoint
self.sampling = self._sampling_config(runtime.SamplingConfig, self.args)
self.eval_args = _make_eval_args(self.args)
self.eval_sampling = self._sampling_config(runtime.SamplingConfig, self.eval_args)
engine_count = self.args.rollout_num_gpus // self.args.rollout_num_gpus_per_engine
self.max_concurrent = self.args.sglang_server_concurrency * engine_count
pool_size = max(1, min(self.max_concurrent, 16))
self.client = _train_client(runtime, self.args, self.model, pool_size)
self.eval_client = _train_client(
runtime,
self.eval_args,
self.model,
pool_size,
router_args=self.args,
)
self.ctx = runtime.ModelContext(client=self.client, model=self.model, sampling=self.sampling)
self.eval_ctx = runtime.ModelContext(client=self.eval_client, model=self.model, sampling=self.eval_sampling)
from miles.utils.misc import load_function
self.dynamic_filter = load_function(self.args.dynamic_sampling_filter_path)
self._tasks = list(self.env.taskset.load())
if self.args.rollout_shuffle:
random.Random(self.args.rollout_seed).shuffle(self._tasks)
if not self._tasks:
raise ValueError("Verifiers taskset selected zero tasks.")
_validate_group_reward_sample_counts(self.args, self._tasks, runtime.discover_decorated)
self._next_train_task_idx = None
self._next_group_index = 0
self._next_sample_index = 0
@staticmethod
def _sampling_config(SamplingConfig, args: Namespace):
data: dict[str, Any] = {
"temperature": args.rollout_temperature,
"top_p": args.rollout_top_p,
"max_tokens": args.rollout_max_response_len,
}
if args.rollout_top_k is not None:
data["top_k"] = args.rollout_top_k
if (min_tokens := getattr(args, "rollout_min_new_tokens", None)) is not None:
data["min_tokens"] = min_tokens
if args.apply_chat_template_kwargs:
data["extra_body"] = {"chat_template_kwargs": args.apply_chat_template_kwargs}
return SamplingConfig.model_validate(data)
def _task(self, index: int):
return self._tasks[index % len(self._tasks)]
async def __call__(self, input: RolloutFnInput) -> RolloutFnOutput:
return await (self._call_eval(input) if input.evaluation else self._call_train(input))
async def _run_task_group(self, task, n: int, semaphore: asyncio.Semaphore, seed_base: int, ctx=None):
runtime = _import_verifiers()
ctx = ctx or self.ctx
episode = self.env.episode(task, ctx, n=n)
if getattr(self.args, "sglang_enable_deterministic_inference", False):
for offset, rollout in enumerate(episode.rollouts):
sampling = ctx.sampling.model_copy(update={"sampling_seed": seed_base + offset})
rollout.ctx = runtime.ModelContext(client=ctx.client, model=self.model, sampling=sampling)
return await episode.run(semaphore)
def _convert_group(self, traces, *, group_index: int, preserve_empty: bool = False):
group = []
complete = True
for trace in traces:
converted = trace_to_samples(
self.args,
trace,
group_index=group_index,
index_start=self._next_sample_index,
)
self._next_sample_index += len(converted)
if not converted:
complete = False
if preserve_empty:
group.append(None)
continue
group.append(converted[0])
return group if complete or preserve_empty else []
async def _apply_miles_rewards(self, group) -> None:
from miles.rollout.rm_hub import async_rm, batched_async_rm
samples = _flatten_samples(group)
if self.args.group_rm:
rewards = await batched_async_rm(self.args, samples)
elif self.args.custom_rm_path is not None or self.args.rm_type:
rewards = await asyncio.gather(*(async_rm(self.args, sample) for sample in samples))
else:
return
if rewards is None or len(rewards) != len(samples):
raise ValueError("Miles reward model returned an unexpected number of rewards.")
for sample, reward in zip(samples, rewards, strict=True):
sample.reward = reward
async def _postprocess_train_samples(self, data, all_data) -> None:
from miles.utils.misc import load_function
if function := load_function(self.args.rollout_sample_filter_path):
function(self.args, data)
if function := load_function(self.args.rollout_all_samples_process_path):
function(self.args, all_data, self.data_source)
await recompute_samples_rollout_logprobs_via_prefill(
self.args,
_flatten_samples(data),
url=_generate_url(self.args),
sampling_params={
"temperature": self.args.rollout_temperature,
"top_p": self.args.rollout_top_p,
"top_k": self.args.rollout_top_k,
"max_new_tokens": self.args.rollout_max_response_len,
},
)
async def _cancel_pending(self, futures: Iterable[asyncio.Task]) -> None:
pending = [future for future in futures if not future.done()]
for future in pending:
future.cancel()
if pending:
try:
from miles.utils.http_utils import post
urls = await _sglang_worker_urls(self.args)
await asyncio.gather(
*(post(f"{url}/abort_request", {"abort_all": True}) for url in urls),
return_exceptions=True,
)
except Exception:
logger.exception("Failed to abort pending Verifiers requests.")
await asyncio.gather(*futures, return_exceptions=True)
async def _call_train(self, input: RolloutFnTrainInput) -> RolloutFnTrainOutput:
from miles.utils import dumper_utils
await dumper_utils.configure_sglang(self.args)
target = self.args.rollout_batch_size
if self._next_train_task_idx is None:
self._next_train_task_idx = input.rollout_id * target
groups = []
all_groups = []
all_traces = []
metrics = MetricGatherer()
semaphore = asyncio.Semaphore(self.max_concurrent)
pending: set[asyncio.Task] = set()
async with self.env.serving():
try:
while len(groups) < target:
while len(groups) + len(pending) < target:
for _ in range(self.args.over_sampling_batch_size):
task_index = self._next_train_task_idx
self._next_train_task_idx += 1
seed = self.args.rollout_seed + task_index * self.args.n_samples_per_prompt
pending.add(
asyncio.create_task(
self._run_task_group(
self._task(task_index),
self.args.n_samples_per_prompt,
semaphore,
seed,
)
)
)
done, pending = await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED)
for future in done:
try:
traces = future.result()
except Exception:
logger.exception("Verifiers episode failed; resampling.")
metrics.on_dynamic_filter_drop(reason="episode_error")
continue
all_traces.extend(traces)
_raise_for_unsupported_trace_errors(traces)
if any(trace.has_error for trace in traces):
metrics.on_dynamic_filter_drop(reason="trace_error")
continue
group = self._convert_group(traces, group_index=self._next_group_index)
self._next_group_index += 1
if len(group) != self.args.n_samples_per_prompt:
metrics.on_dynamic_filter_drop(reason="empty_trace")
continue
await self._apply_miles_rewards(group)
all_groups.append(group)
result = call_dynamic_filter(
self.dynamic_filter,
self.args,
_flatten_samples(group),
)
if not result.keep:
metrics.on_dynamic_filter_drop(reason=result.reason)
continue
if len(groups) < target:
groups.append(group)
finally:
await self._cancel_pending(pending)
groups.sort(key=lambda group: _flatten_samples(group)[0].index)
all_groups.sort(key=lambda group: _flatten_samples(group)[0].index)
await self._postprocess_train_samples(groups, all_groups)
output_metrics = metrics.collect()
output_metrics.update(_trace_metrics(all_traces))
return RolloutFnTrainOutput(samples=groups, metrics=output_metrics)
async def _call_eval(self, input: RolloutFnEvalInput) -> RolloutFnEvalOutput:
assert not self.args.group_rm, "Group RM is not supported for eval rollout"
from miles.utils import dumper_utils
await dumper_utils.configure_sglang(self.args)
semaphore = asyncio.Semaphore(self.max_concurrent)
async with self.env.serving():
futures = [
asyncio.create_task(
self._run_task_group(
task,
self.args.n_samples_per_eval_prompt,
semaphore,
self.args.rollout_seed + index * self.args.n_samples_per_eval_prompt,
self.eval_ctx,
)
)
for index, task in enumerate(self._tasks)
]
try:
trace_groups = await asyncio.gather(*futures)
finally:
await self._cancel_pending(futures)
samples = []
rewards = []
truncated = []
all_traces = []
reward_key = self.args.eval_reward_key or self.args.reward_key
use_miles_rewards = bool(self.args.custom_rm_path is not None or self.args.rm_type)
for group_index, traces in enumerate(trace_groups):
all_traces.extend(traces)
_raise_for_unsupported_trace_errors(traces)
group = self._convert_group(traces, group_index=group_index, preserve_empty=True)
trainable = [value for value in group if value is not None]
await self._apply_miles_rewards(trainable)
samples.extend(_flatten_samples(trainable))
for trace, value in zip(traces, group, strict=True):
if use_miles_rewards and value is not None:
sample_rewards = [sample.get_reward_value(self.eval_args) for sample in _flatten_samples([value])]
reward = sum(sample_rewards) / len(sample_rewards)
else:
reward = trace.reward if use_miles_rewards else _trace_eval_reward(trace, reward_key)
rewards.append(reward)
truncated.append(trace.is_truncated)
return RolloutFnEvalOutput(
data={
self.config.env_id
or "verifiers": {
"rewards": rewards,
"truncated": truncated,
"samples": samples,
}
},
metrics=_trace_metrics(all_traces),
)
_LEGACY_INSTANCES: dict[tuple[int, int, bool], VerifiersRolloutFn] = {}
def generate_rollout(
args: Namespace,
rollout_id: int,
data_source: Any,
evaluation: bool = False,
) -> RolloutFnTrainOutput | RolloutFnEvalOutput:
"""Legacy Miles entrypoint backed by one persistent rollout adapter."""
from miles.utils.async_utils import run
# The refactored rollout manager constructs separate train and eval adapters.
# Preserve that lifecycle under the legacy function interface as well.
key = (id(args), id(data_source), evaluation)
adapter = _LEGACY_INSTANCES.get(key)
if adapter is None:
adapter = VerifiersRolloutFn(RolloutFnConstructorInput(args=args, data_source=data_source))
_LEGACY_INSTANCES[key] = adapter
input = RolloutFnEvalInput(rollout_id) if evaluation else RolloutFnTrainInput(rollout_id)
return run(adapter(input))
+146
View File
@@ -0,0 +1,146 @@
import math
import os
import shutil
import sys
from collections import Counter
from pathlib import Path
import pytest
import torch
from tests.ci.ci_register import register_cuda_ci
import miles.utils.external_utils.command_utils as U
register_cuda_ci(est_time=900, suite="stage-c-2-gpu-h200", labels=["long"])
MODEL_NAME = "Qwen3-0.6B"
MODEL_TYPE = "qwen3-0.6B"
NUM_GPUS = 2
MODEL_DIR = Path(os.environ.get("MILES_E2E_MODEL_DIR", "/root/models"))
MEGATRON_PATH = Path(os.environ.get("MILES_E2E_MEGATRON_PATH", "/root/Megatron-LM"))
RUN_DIR = Path(os.environ.get("MILES_E2E_RUN_DIR", "/tmp/miles-verifiers-e2e"))
VERIFIERS_DIR = Path("/tmp/verifiers-v0.2.0")
ADAPTER_DIR = Path(U.repo_base_dir) / "examples" / "experimental" / "verifiers"
def prepare():
U.exec_command(f"mkdir -p {MODEL_DIR} {RUN_DIR}")
U.exec_command(f"hf download Qwen/{MODEL_NAME} --local-dir {MODEL_DIR}/{MODEL_NAME}")
U.exec_command(
f"{sys.executable} -m pip install -r {U.repo_base_dir}/examples/experimental/verifiers/requirements.txt"
)
U.exec_command("uv tool install 'prime==0.6.19'")
if not VERIFIERS_DIR.exists():
U.exec_command(
f"git clone --depth 1 --branch v0.2.0 "
f"https://github.com/PrimeIntellect-ai/verifiers.git {VERIFIERS_DIR}"
)
shutil.copytree(
VERIFIERS_DIR / "environments" / "code_golf_v1",
RUN_DIR / "environments" / "code_golf_v1",
dirs_exist_ok=True,
)
U.exec_command(f"cd {RUN_DIR} && prime --plain env install code-golf-v1")
U.convert_checkpoint(
model_name=MODEL_NAME,
megatron_model_type=MODEL_TYPE,
num_gpus_per_node=NUM_GPUS,
dir_dst=str(MODEL_DIR),
hf_checkpoint=str(MODEL_DIR / MODEL_NAME),
megatron_path=str(MEGATRON_PATH),
)
def execute():
config_path = RUN_DIR / "code-golf.toml"
config_path.write_text('[taskset]\nid = "code-golf-v1"\n')
dump_dir = RUN_DIR / "dump"
train_args = " ".join(
[
f"--hf-checkpoint {MODEL_DIR}/{MODEL_NAME}",
"--sglang-tokenizer-path Qwen/Qwen3-0.6B",
f"--ref-load {MODEL_DIR}/{MODEL_NAME}_torch_dist",
"--rollout-function-path verifiers_rollout.generate_rollout",
"--disable-rollout-global-dataset",
"--num-rollout 1",
"--rollout-batch-size 3",
"--n-samples-per-prompt 4",
"--over-sampling-batch-size 3",
"--rollout-max-response-len 512",
"--rollout-max-context-len 2048",
"--rollout-temperature 0.8",
"--global-batch-size 12",
"--balance-data",
"--advantage-estimator grpo",
"--entropy-coef 0.0",
"--eps-clip 0.2",
"--eps-clip-high 0.28",
"--optimizer adam",
"--lr 1e-6",
"--lr-decay-style constant",
"--weight-decay 0.1",
"--adam-beta1 0.9",
"--adam-beta2 0.98",
"--no-gradient-accumulation-fusion",
"--rollout-num-gpus-per-engine 1",
"--sglang-mem-fraction-static 0.6",
"--sglang-enable-metrics",
"--tensor-model-parallel-size 1",
"--pipeline-model-parallel-size 1",
"--context-parallel-size 1",
"--use-dynamic-batch-size",
"--max-tokens-per-gpu 4096",
"--actor-num-nodes 1",
f"--actor-num-gpus-per-node {NUM_GPUS}",
"--colocate",
f"--dump-details {dump_dir}",
U.get_default_wandb_args(__file__),
]
)
U.execute_train(
train_args=train_args,
num_gpus_per_node=NUM_GPUS,
megatron_model_type=MODEL_TYPE,
extra_env_vars={
"MILES_EXPERIMENTAL_ROLLOUT_REFACTOR": "0",
"PYTHONPATH": f"{MEGATRON_PATH}:{ADAPTER_DIR}:{U.repo_base_dir}",
"VERIFIERS_CONFIG": str(config_path),
},
megatron_path=str(MEGATRON_PATH),
)
verify(dump_dir)
def verify(dump_dir: Path):
samples = torch.load(dump_dir / "rollout_data" / "0.pt", weights_only=False)["samples"]
assert len(samples) == 12
assert Counter(sample["group_index"] for sample in samples) == {0: 4, 1: 4, 2: 4}
assert {sample["status"] for sample in samples} <= {"completed", "truncated"}
assert any(sample["status"] == "completed" for sample in samples)
assert all(math.isfinite(sample["reward"]) for sample in samples)
assert all(
len(sample["rollout_log_probs"]) == len(sample["loss_mask"]) == sample["response_length"] for sample in samples
)
assert all(math.isfinite(value) for sample in samples for value in sample["rollout_log_probs"])
for sample in samples:
metadata = sample["metadata"]["verifiers"]
assert set(metadata["rewards"]) == {"correct", "fastest", "most_concise"}
assert metadata["metrics"]["passed"] in {0.0, 1.0}
assert math.isfinite(metadata["metrics"]["latency"])
assert "error" not in metadata
assert any(sample["metadata"]["verifiers"]["metrics"]["passed"] for sample in samples)
for group_index in range(3):
group = [sample for sample in samples if sample["group_index"] == group_index]
assert sum(sample["metadata"]["verifiers"]["rewards"]["fastest"] for sample in group) == pytest.approx(0.5)
if __name__ == "__main__":
prepare()
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
os.environ.pop(proxy_var, None)
execute()
@@ -0,0 +1,698 @@
import asyncio
import sys
from argparse import Namespace
from contextlib import asynccontextmanager
from types import SimpleNamespace
import pytest
from packaging.version import Version
if sys.version_info < (3, 11):
pytest.skip("Verifiers requires Python 3.11+", allow_module_level=True)
from examples.experimental.verifiers.verifiers_rollout import (
MilesSGLangTransport,
VerifiersRolloutFn,
_check_version,
_config_path,
_finish_reason,
_load_config_data,
_make_eval_args,
_raise_for_unsupported_trace_errors,
_renderer_identity,
_trace_eval_reward,
_train_client,
_validate_args,
_validate_group_reward_sample_counts,
trace_to_sample,
trace_to_samples,
)
from miles.utils.types import Sample
def _args(**overrides) -> Namespace:
values = {
"lora_adapter_path": None,
"lora_rank": 0,
"reward_key": None,
"rollout_max_context_len": 64,
"rollout_max_prompt_len": None,
"rollout_max_response_len": 8,
"rollout_skip_special_tokens": True,
"rollout_stop": None,
"rollout_stop_token_ids": None,
"sglang_model_routers": None,
"sglang_router_ip": "127.0.0.1",
"sglang_router_policy": "round_robin",
"sglang_router_port": 30000,
"sglang_tokenizer_path": None,
}
values.update(overrides)
return Namespace(**values)
def _branch(*, index=0, token_ids=None, sampled_mask=None, logprobs=None):
return SimpleNamespace(
index=index,
token_ids=[10, 11, 20, 21, 22] if token_ids is None else token_ids,
sampled_mask=[False, False, True, False, True] if sampled_mask is None else sampled_mask,
logprobs=[0.0, 0.0, -0.1, 0.0, -0.2] if logprobs is None else logprobs,
)
def _trace(**overrides):
values = {
"id": "trace-1",
"branches": [_branch()],
"task": SimpleNamespace(data=SimpleNamespace(prompt="solve this", idx="task-1")),
"rewards": {"score": 1.25, "bonus": 0.75},
"metrics": {"turns": 2.0},
"stop_condition": "done",
"error": None,
"has_error": False,
"is_truncated": False,
"reward": 2.0,
"last_reply": "answer",
"num_turns": 2,
}
values.update(overrides)
return SimpleNamespace(**values)
def _wiring_args(**overrides) -> Namespace:
values = {
"rollout_global_dataset": False,
"partial_rollout": False,
"multimodal_keys": None,
"chat_template_path": None,
"use_opd": False,
"use_rollout_routing_replay": False,
"use_rollout_indexer_replay": False,
}
values.update(overrides)
return Namespace(**values)
def test_supported_wiring_passes_validation():
_validate_args(_wiring_args())
@pytest.mark.parametrize(
("overrides", "message"),
[
({"rollout_global_dataset": True}, "--disable-rollout-global-dataset"),
({"partial_rollout": True}, "cannot be resumed"),
({"multimodal_keys": {"image": "image"}}, "text-only renderer inputs"),
({"chat_template_path": "/tmp/custom.jinja"}, "custom\n?\s*Jinja template"),
({"use_opd": True}, "--use-opd"),
({"use_rollout_routing_replay": True}, "--use-rollout-routing-replay"),
({"use_rollout_indexer_replay": True}, "--use-rollout-indexer-replay"),
],
)
def test_unsupported_wiring_fails_before_any_episode_runs(overrides, message):
with pytest.raises(ValueError, match=message):
_validate_args(_wiring_args(**overrides))
def test_config_path_comes_from_the_environment(monkeypatch):
monkeypatch.setenv("VERIFIERS_CONFIG", "/tmp/vf.toml")
assert _config_path() == "/tmp/vf.toml"
def test_missing_config_env_var_names_the_variable(monkeypatch):
monkeypatch.delenv("VERIFIERS_CONFIG", raising=False)
with pytest.raises(ValueError, match="VERIFIERS_CONFIG"):
_config_path()
def test_config_loader_uses_verifiers_toml_format(tmp_path):
path = tmp_path / "verifiers.toml"
path.write_text('[taskset]\nid = "gsm8k-v1"\n')
assert _load_config_data(path) == {"taskset": {"id": "gsm8k-v1"}}
def test_config_loader_rejects_non_toml_formats(tmp_path):
path = tmp_path / "verifiers.yaml"
path.write_text("taskset: gsm8k-v1\n")
with pytest.raises(ValueError, match="Verifiers TOML"):
_load_config_data(path)
def test_verifiers_0_2_1_is_rejected_with_compatibility_reason():
with pytest.raises(RuntimeError, match="SGLang 0.5.15 pins OpenAI"):
_check_version("verifiers", "0.2.1", Version("0.2.0"), Version("0.2.1"))
def test_transport_does_not_treat_aborted_generation_as_complete():
with pytest.raises(RuntimeError, match="aborted"):
_finish_reason({"meta_info": {"finish_reason": {"type": "abort"}}})
@pytest.mark.parametrize(
("checkpoint", "expected"),
[
("/models/Qwen3-4B-Instruct-2507", "Qwen/Qwen3-4B-Instruct-2507"),
(
"/cache/models--Qwen--Qwen3-4B-Instruct-2507/snapshots/revision",
"Qwen/Qwen3-4B-Instruct-2507",
),
("/models/private-finetune", None),
],
)
def test_renderer_identity_is_inferred_from_standard_checkpoint_paths(checkpoint, expected):
pytest.importorskip("renderers", minversion="0.1.8")
assert _renderer_identity(checkpoint) == expected
def test_train_client_uses_local_tokenizer_with_inferred_renderer_identity(monkeypatch):
renderers = pytest.importorskip("renderers", minversion="0.1.8")
checkpoint = "/cache/models--Qwen--Qwen3-4B-Instruct-2507/snapshots/revision"
seen = {}
class BaseTrainClient:
def __init__(self, openai, pool_size, config=None, renderer_model_name=None):
self.openai = openai
self.pool_size = pool_size
self.config = config
self.renderer_model_name = renderer_model_name
self._pool = None
class RendererPool:
def __init__(self, factory, size):
seen["size"] = size
seen["renderer"] = factory()
def load_tokenizer(source):
seen["source"] = source
return SimpleNamespace(name_or_path=source)
def create_renderer(tokenizer, config, *, chat_template_kwargs=None):
seen["identity"] = tokenizer.name_or_path
seen["config"] = config
seen["kwargs"] = chat_template_kwargs
return "renderer"
monkeypatch.setattr(renderers, "RendererPool", RendererPool)
monkeypatch.setattr(renderers, "create_renderer", create_renderer)
monkeypatch.setattr("renderers.base.load_tokenizer", load_tokenizer)
runtime = SimpleNamespace(TrainClient=BaseTrainClient)
args = _args(sglang_tokenizer_path="/models/custom-tokenizer")
client = _train_client(runtime, args, checkpoint, pool_size=3)
pool = client._renderer_pool(checkpoint, chat_template_kwargs={"enable_thinking": False})
assert isinstance(pool, RendererPool)
assert seen == {
"size": 3,
"source": "/models/custom-tokenizer",
"identity": "Qwen/Qwen3-4B-Instruct-2507",
"config": None,
"kwargs": {"enable_thinking": False},
"renderer": "renderer",
}
@pytest.mark.asyncio
async def test_train_client_reports_unsupported_tool_renderer_as_configuration_error():
pytest.importorskip("renderers", minversion="0.1.8")
class ProviderError(Exception):
def __init__(self, message, *, status_code):
super().__init__(message)
self.status_code = status_code
class BaseTrainClient:
def __init__(self, openai, pool_size, config=None, renderer_model_name=None):
pass
async def get_response(self, *args, **kwargs):
raise ValueError("RendererPool does not support tools.")
runtime = SimpleNamespace(ProviderError=ProviderError, TrainClient=BaseTrainClient)
client = _train_client(runtime, _args(), "/models/private-finetune", pool_size=1)
with pytest.raises(ProviderError, match="--sglang-tokenizer-path") as error:
await client.get_response()
assert error.value.status_code == 400
@pytest.mark.asyncio
async def test_train_client_reports_unsupported_dialect_as_configuration_error():
pytest.importorskip("renderers", minversion="0.1.8")
class ProviderError(Exception):
def __init__(self, message, *, status_code):
super().__init__(message)
self.status_code = status_code
class BaseTrainClient:
def __init__(self, openai, pool_size, config=None, renderer_model_name=None):
pass
async def get_response(self, *args, **kwargs):
raise NotImplementedError("only the chat-completions dialect is supported")
runtime = SimpleNamespace(ProviderError=ProviderError, TrainClient=BaseTrainClient)
client = _train_client(runtime, _args(), "/models/test", pool_size=1)
with pytest.raises(ProviderError, match="does not support this request") as error:
await client.get_response()
assert error.value.status_code == 400
@pytest.mark.parametrize(("method", "message"), [("relay", "streaming"), ("relay_aux", "auxiliary")])
@pytest.mark.asyncio
async def test_train_client_reports_unsupported_relay_paths_as_configuration_errors(method, message):
pytest.importorskip("renderers", minversion="0.1.8")
class ProviderError(Exception):
def __init__(self, text, *, status_code):
super().__init__(text)
self.status_code = status_code
class BaseTrainClient:
def __init__(self, openai, pool_size, config=None, renderer_model_name=None):
pass
runtime = SimpleNamespace(ProviderError=ProviderError, TrainClient=BaseTrainClient)
client = _train_client(runtime, _args(), "/models/test", pool_size=1)
with pytest.raises(ProviderError, match=message) as error:
await getattr(client, method)()
assert error.value.status_code == 400
def test_trace_to_sample_preserves_training_fields_and_verifiers_reward():
sample = trace_to_sample(_args(), _trace(), group_index=3, index=9)
assert sample.group_index == 3
assert sample.index == 9
assert sample.prompt == "solve this"
assert sample.tokens == [10, 11, 20, 21, 22]
assert sample.response == "answer"
assert sample.response_length == 3
assert sample.loss_mask == [1, 0, 1]
assert sample.rollout_log_probs == [-0.1, 0.0, -0.2]
assert sample.reward == 2.0
assert sample.routing_key == "trace-1"
assert sample.status == Sample.Status.COMPLETED
assert sample.metadata["verifiers"]["task_index"] == "task-1"
sample.validate()
def test_trace_to_sample_preserves_named_rewards_when_reward_key_is_set():
args = _args(reward_key="score")
sample = trace_to_sample(args, _trace(), group_index=0, index=0)
assert sample.reward == {"score": 1.25, "bonus": 0.75, "reward": 2.0}
assert sample.get_reward_value(args) == 1.25
def test_trace_to_sample_serializes_structured_prompt_messages():
class Message:
def model_dump(self, **kwargs):
assert kwargs == {"mode": "json", "exclude_none": True}
return {"role": "user", "content": "solve this"}
trace = _trace(task=SimpleNamespace(data=SimpleNamespace(prompt=[Message()], idx="task-1")))
sample = trace_to_sample(_args(), trace, group_index=0, index=0)
assert sample.prompt == [{"role": "user", "content": "solve this"}]
def test_trace_to_sample_marks_error_before_truncation():
error = SimpleNamespace(model_dump=lambda **_kwargs: {"type": "ProviderError"})
sample = trace_to_sample(
_args(),
_trace(error=error, has_error=True, is_truncated=True),
group_index=0,
index=0,
)
assert sample.status == Sample.Status.FAILED
assert sample.metadata["verifiers"]["error"] == {"type": "ProviderError"}
def test_failed_eval_trace_with_missing_named_reward_returns_none():
trace = _trace(has_error=True, rewards={})
assert _trace_eval_reward(trace, "score") is None
def test_successful_eval_trace_requires_configured_named_reward():
with pytest.raises(KeyError, match="score"):
_trace_eval_reward(_trace(rewards={}), "score")
def test_unsupported_trace_error_is_not_resampled_forever():
error = SimpleNamespace(message="Miles' Verifiers adapter does not support this request: ResponsesDialect")
with pytest.raises(RuntimeError, match="ResponsesDialect"):
_raise_for_unsupported_trace_errors([_trace(error=error, has_error=True)])
def test_graph_branches_fail_before_miles_can_corrupt_trace_groups():
trace = _trace(branches=[_branch(index=0), _branch(index=1)])
with pytest.raises(NotImplementedError, match="multiple graph branches"):
trace_to_samples(_args(), trace, group_index=4, index_start=10)
def test_convert_group_uses_standard_miles_group_shape():
from miles.ray.rollout.rollout_data_conversion import postprocess_rollout_data
adapter = object.__new__(VerifiersRolloutFn)
adapter.args = _args()
adapter._next_sample_index = 0
group = adapter._convert_group(
[_trace(id="first"), _trace(id="second")],
group_index=2,
)
assert len(group) == 2
assert all(isinstance(sample, Sample) for sample in group)
args = SimpleNamespace(
disable_rollout_trim_samples=True,
global_batch_size=1,
use_dynamic_global_batch_size=False,
)
flattened, _ = postprocess_rollout_data(args, [group], train_parallel_config={"dp_size": 1})
assert [sample.routing_key for sample in flattened] == ["first", "second"]
def test_standard_dynamic_filter_accepts_converted_branch_group():
from miles.rollout.filter_hub.dynamic_sampling_filters import check_reward_nonzero_std
adapter = object.__new__(VerifiersRolloutFn)
adapter.args = _args()
adapter._next_sample_index = 0
group = adapter._convert_group(
[_trace(id="low", reward=0.0), _trace(id="high", reward=1.0)],
group_index=0,
)
assert check_reward_nonzero_std(adapter.args, group).keep
@pytest.mark.parametrize(
("overrides", "option"),
[
(
{"eval_interval": None, "n_samples_per_prompt": 1, "n_samples_per_eval_prompt": 1},
"--n-samples-per-prompt",
),
(
{"eval_interval": 1, "n_samples_per_prompt": 2, "n_samples_per_eval_prompt": 1},
"--n-samples-per-eval-prompt",
),
],
)
def test_group_reward_tasks_require_multiple_rollouts(overrides, option):
args = _args(**overrides)
tasks = [object()]
with pytest.raises(ValueError, match=option):
_validate_group_reward_sample_counts(args, tasks, lambda _task, _kind: [object()])
def test_group_reward_eval_count_is_ignored_when_eval_is_disabled():
args = _args(
eval_interval=None,
n_samples_per_prompt=2,
n_samples_per_eval_prompt=1,
)
_validate_group_reward_sample_counts(args, [object()], lambda _task, _kind: [object()])
def test_group_reward_train_count_is_ignored_for_eval_only_runs():
args = _args(
num_rollout=0,
eval_interval=1,
n_samples_per_prompt=1,
n_samples_per_eval_prompt=2,
)
_validate_group_reward_sample_counts(args, [object()], lambda _task, _kind: [object()])
@pytest.mark.asyncio
async def test_verifiers_episode_owns_group_reward_computation():
pytest.importorskip("verifiers", minversion="0.2.0")
pytest.importorskip("renderers", minversion="0.1.8")
traces = [_trace(id="a", reward=0.0), _trace(id="b", reward=0.0)]
class Episode:
rollouts = []
async def run(self, semaphore):
assert semaphore is not None
traces[0].reward = -1.0
traces[1].reward = 1.0
return traces
class Environment:
def episode(self, task, ctx, n):
assert task == "task"
assert ctx == "ctx"
assert n == 2
return Episode()
adapter = object.__new__(VerifiersRolloutFn)
adapter.args = _args(sglang_enable_deterministic_inference=False)
adapter.env = Environment()
adapter.ctx = "ctx"
result = await adapter._run_task_group("task", 2, asyncio.Semaphore(2), seed_base=0)
assert [trace.reward for trace in result] == [-1.0, 1.0]
def test_sampling_config_preserves_miles_minimum_tokens():
class SamplingConfig:
@staticmethod
def model_validate(data):
return data
config = VerifiersRolloutFn._sampling_config(
SamplingConfig,
_args(
apply_chat_template_kwargs={},
rollout_min_new_tokens=3,
rollout_temperature=0.7,
rollout_top_k=20,
rollout_top_p=0.9,
),
)
assert config["min_tokens"] == 3
def test_eval_args_clear_training_prompt_cap_and_preserve_other_fallbacks():
args = _args(
eval_max_context_len=128,
eval_max_prompt_len=None,
eval_max_response_len=None,
eval_min_new_tokens=None,
eval_reward_key=None,
eval_temperature=None,
eval_top_k=None,
eval_top_p=None,
reward_key="score",
rollout_max_context_len=64,
rollout_max_prompt_len=32,
rollout_max_response_len=8,
)
eval_args = _make_eval_args(args)
assert eval_args.rollout_max_context_len == 128
assert eval_args.rollout_max_prompt_len is None
assert eval_args.rollout_max_response_len == 8
assert eval_args.reward_key == "score"
@pytest.mark.asyncio
async def test_transport_translates_renderer_request_to_sglang(monkeypatch):
requests = []
async def fake_post(url, payload, headers=None):
requests.append((url, payload, headers))
return {
"request_id": "request-id",
"meta_info": {
"completion_tokens": 2,
"finish_reason": {"type": "stop"},
"output_token_logprobs": [[-0.1, 20], [-0.2, 21]],
},
}
monkeypatch.setattr("miles.utils.http_utils.post", fake_post)
transport = MilesSGLangTransport(_args(sglang_router_policy="manual"))
response = await transport.post(
"http://127.0.0.1:30000/inference/v1/generate",
body={
"model": "test/model",
"token_ids": [10, 11],
"sampling_params": {"temperature": 0.2, "max_tokens": 2, "stop_token_ids": [99], "logprobs": 1},
},
options={"headers": {"X-Session-ID": "trace-id"}},
)
assert response.json()["choices"][0]["token_ids"] == [20, 21]
assert requests == [
(
"http://127.0.0.1:30000/generate",
{
"input_ids": [10, 11],
"sampling_params": {
"temperature": 0.2,
"stop_token_ids": [99],
"max_new_tokens": 2,
"skip_special_tokens": True,
"no_stop_trim": True,
"spaces_between_special_tokens": False,
"n": 1,
},
"return_logprob": True,
},
{"X-SMG-Routing-Key": "trace-id"},
)
]
@pytest.mark.asyncio
async def test_transport_rejects_multimodal_features():
transport = MilesSGLangTransport(_args())
with pytest.raises(NotImplementedError, match="multimodal"):
await transport.post(
"http://127.0.0.1:30000/inference/v1/generate",
body={"token_ids": [1], "sampling_params": {}, "features": {}},
)
@pytest.mark.asyncio
async def test_transport_bounds_seen_sessions(monkeypatch):
async def fake_post(_url, _payload, headers=None):
return {
"meta_info": {
"completion_tokens": 1,
"finish_reason": {"type": "stop"},
"output_token_logprobs": [[-0.1, 20]],
}
}
monkeypatch.setattr("miles.utils.http_utils.post", fake_post)
transport = MilesSGLangTransport(_args())
transport._session_cache_size = 2
body = {"token_ids": [10], "sampling_params": {}}
for session_id in ("one", "two", "three"):
await transport.post(
"http://127.0.0.1:30000/inference/v1/generate",
body=body,
options={"headers": {"X-Session-ID": session_id}},
)
assert list(transport._seen_sessions) == ["two", "three"]
def test_transport_resolves_router_address_lazily():
args = _args(sglang_router_ip=None, sglang_router_port=None)
transport = MilesSGLangTransport(args)
args.sglang_router_ip = "10.0.0.4"
args.sglang_router_port = 3210
assert transport.base_url == "http://10.0.0.4:3210/v1"
def test_transport_uses_default_model_router():
args = _args(
sglang_model_routers={
"default": ("10.0.0.5", 3211),
"ref": ("10.0.0.6", 3212),
}
)
assert MilesSGLangTransport(args).base_url == "http://10.0.0.5:3211/v1"
def test_eval_transport_keeps_live_router_args():
rollout_args = _args(sglang_router_ip=None, sglang_router_port=None)
eval_args = _args(sglang_router_ip=None, sglang_router_port=None)
transport = MilesSGLangTransport(eval_args, router_args=rollout_args)
rollout_args.sglang_model_routers = {"default": ("10.0.0.7", 3213)}
assert transport.base_url == "http://10.0.0.7:3213/v1"
@pytest.mark.asyncio
async def test_eval_rejects_miles_group_reward_model():
adapter = object.__new__(VerifiersRolloutFn)
adapter.args = _args(group_rm=True)
with pytest.raises(AssertionError, match="Group RM is not supported for eval rollout"):
await adapter._call_eval(SimpleNamespace(rollout_id=0))
@pytest.mark.asyncio
async def test_eval_extracts_structured_miles_reward(monkeypatch):
import miles.utils as miles_utils
@asynccontextmanager
async def serving():
yield
async def configure_sglang(_args):
return None
async def run_group(*_args, **_kwargs):
return [_trace(rewards={"bonus": 0.25})]
async def apply_reward(group):
group[0].reward = {"score": 0.75, "details": "ok"}
dumper_utils = SimpleNamespace(configure_sglang=configure_sglang)
monkeypatch.setitem(sys.modules, "miles.utils.dumper_utils", dumper_utils)
monkeypatch.setattr(miles_utils, "dumper_utils", dumper_utils, raising=False)
adapter = object.__new__(VerifiersRolloutFn)
adapter.args = _args(
custom_rm_path="tests.fake_reward",
eval_reward_key="score",
group_rm=False,
n_samples_per_eval_prompt=1,
rm_type=None,
rollout_batch_size=1,
rollout_seed=1,
)
adapter.eval_args = _args(reward_key="score")
adapter.config = SimpleNamespace(env_id="test-env")
adapter.env = SimpleNamespace(serving=serving)
adapter.eval_ctx = object()
adapter.max_concurrent = 1
adapter.model = "test/model"
adapter._tasks = [object()]
adapter._next_sample_index = 0
adapter._run_task_group = run_group
adapter._apply_miles_rewards = apply_reward
output = await adapter._call_eval(SimpleNamespace(rollout_id=0))
assert output.data["test-env"]["rewards"] == [0.75]
@@ -0,0 +1,122 @@
import sys
from argparse import Namespace
from types import SimpleNamespace
import pytest
if sys.version_info < (3, 11):
pytest.skip("Verifiers requires Python 3.11+", allow_module_level=True)
pytest.importorskip("verifiers", minversion="0.2.0")
pytest.importorskip("renderers", minversion="0.1.8")
from examples.experimental.verifiers.verifiers_rollout import MilesSGLangTransport
from verifiers.v1.clients.train import TrainClient
from verifiers.v1.dialects import ChatDialect, ResponsesDialect
from verifiers.v1.env import EnvConfig, Environment
from verifiers.v1.types import SamplingConfig
def _args(**overrides):
values = {
"lora_adapter_path": None,
"lora_rank": 0,
"rollout_max_context_len": 64,
"rollout_max_prompt_len": None,
"rollout_max_response_len": 8,
"rollout_skip_special_tokens": True,
"rollout_stop": None,
"rollout_stop_token_ids": None,
"sglang_model_routers": None,
"sglang_router_ip": "127.0.0.1",
"sglang_router_policy": "round_robin",
"sglang_router_port": 30000,
"sglang_tokenizer_path": None,
}
values.update(overrides)
return Namespace(**values)
def test_minimal_env_config_uses_the_v1_environment_contract():
config = EnvConfig.model_validate({"taskset": {"id": "harbor"}})
environment = Environment(config)
assert config.is_legacy is False
assert config.env_id == "harbor"
assert type(environment.taskset).__name__ == "HarborTaskset"
class _Rendered:
token_ids = [10, 11]
multi_modal_data = None
is_content = [True, True]
@staticmethod
def message_token_spans():
return [(0, 2)]
class _Renderer:
supports_tools = True
def render(self, messages, *, tools, add_generation_prompt):
assert messages == [{"role": "user", "content": "question"}]
assert tools is None
assert add_generation_prompt is True
return _Rendered()
@staticmethod
def get_stop_token_ids():
return [99]
@staticmethod
def parse_response(token_ids, *, tools):
assert token_ids == [20, 21]
assert tools is None
return SimpleNamespace(content="answer", reasoning_content=None, tool_calls=[])
@pytest.mark.asyncio
async def test_published_train_client_runs_through_miles_transport(monkeypatch):
async def fake_post(_url, _payload, headers=None):
assert headers is None
return {
"request_id": "request-id",
"meta_info": {
"completion_tokens": 2,
"finish_reason": {"type": "stop"},
"output_token_logprobs": [[-0.1, 20], [-0.2, 21]],
},
}
monkeypatch.setattr("miles.utils.http_utils.post", fake_post)
client = TrainClient(MilesSGLangTransport(_args()), renderer_model_name="test/model")
client._pool = _Renderer()
response = await client.get_response(
ChatDialect(),
{"messages": [{"role": "user", "content": "question"}]},
"test/model",
SamplingConfig(temperature=0.2, max_tokens=2),
session_id="trace-id",
)
assert response.message.content == "answer"
assert response.tokens.prompt_ids == [10, 11]
assert response.tokens.completion_ids == [20, 21]
assert response.tokens.completion_logprobs == [-0.1, -0.2]
ChatDialect().validate_response(response.raw)
@pytest.mark.asyncio
async def test_published_train_client_rejects_non_chat_dialects():
client = TrainClient(MilesSGLangTransport(_args()), renderer_model_name="test/model")
with pytest.raises(NotImplementedError, match="chat-completions dialect"):
await client.get_response(
ResponsesDialect(),
{"input": "question"},
"test/model",
SamplingConfig(max_tokens=2),
)