mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
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:
co-authored by
Tao Lin
Claude Fable 5
parent
03491d0431
commit
d2010d2929
+2
-1
@@ -182,7 +182,8 @@
|
||||
"pages": [
|
||||
"user-guide/harbor",
|
||||
"user-guide/openenv",
|
||||
"user-guide/nemo-gym"
|
||||
"user-guide/nemo-gym",
|
||||
"user-guide/verifiers"
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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))
|
||||
@@ -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),
|
||||
)
|
||||
Reference in New Issue
Block a user