Add a solver-verifier example on gsm8k (#2645)

This commit is contained in:
fzyzcjy
2026-09-26 19:08:04 +08:00
committed by GitHub
parent f7c5dee725
commit 8a12b98547
9 changed files with 857 additions and 0 deletions
@@ -46,6 +46,7 @@
"run-ci-miles-plugin",
"run-ci-model-scripts",
"run-ci-mooncake",
"run-ci-multi-policy",
"run-ci-precision",
"run-ci-qwen35",
"run-ci-replay",
View File
+132
View File
@@ -0,0 +1,132 @@
# See tests/e2e/short/test_multi_policy_solver_verifier_gsm8k.py for an end-to-end run of this example.
import dataclasses
import re
from enum import Enum
from miles.backends.megatron_utils.megatron_config import resolve_megatron_config
from miles.rollout.base_types import GenerateFnInput, GenerateFnOutput
from miles.rollout.generate_hub.single_turn import generate as single_turn_generate
from miles.utils.types import Sample
_AGREE_MARKER = "AGREE"
_WRONG_MARKER = "WRONG"
_VERDICT_PREFIX = "VERDICT:"
_VERIFIER_PROMPT_TEMPLATE = (
"Another model was asked to solve a math problem, and you must check its work.\n\n"
"Question:\n{question}\n\n"
"Proposed solution:\n{solver_response}\n\n"
f"Reason about the proposed solution, then write exactly one verdict line of its own: "
f"'{_VERDICT_PREFIX} {_AGREE_MARKER}' if the proposed final answer is correct, or "
f"'{_VERDICT_PREFIX} {_WRONG_MARKER}' if it is not. After a {_WRONG_MARKER} verdict, "
"write your own final answer on a last line as '#### <answer>'."
)
_VERDICT_PATTERN = re.compile(rf"^{_VERDICT_PREFIX}\s*({_AGREE_MARKER}|{_WRONG_MARKER})\s*$", re.MULTILINE)
_MARKED_ANSWER_PATTERN = re.compile(r"####\s*([^\n]+)")
_NUMBER_PATTERN = re.compile(r"-?\d+(?:[\d,]*\d)?(?:\.\d+)?")
class _Verdict(Enum):
AGREE = _AGREE_MARKER
WRONG = _WRONG_MARKER
async def generate(input: GenerateFnInput) -> GenerateFnOutput:
args = input.args
model_ids = resolve_megatron_config(args).model_ids
assert len(model_ids) == 2, (
f"examples/multi_policy/solver_verifier.py pairs one solver policy with one verifier policy, but "
f"--megatron-config names {model_ids}"
)
solver_model_id, verifier_model_id = model_ids
solver_output = await single_turn_generate(input, url=_compute_router_url(args, model_id=solver_model_id))
solver_sample = solver_output.samples
assert isinstance(solver_sample, Sample), f"{solver_sample=}"
assert solver_sample.status != Sample.Status.ABORTED
verifier_sample = _build_verifier_sample(solver_sample)
verifier_output = await single_turn_generate(
dataclasses.replace(input, sample=verifier_sample),
url=_compute_router_url(args, model_id=verifier_model_id),
)
verifier_sample = verifier_output.samples
assert isinstance(verifier_sample, Sample), f"{verifier_sample=}"
ground_truth = _extract_answer(solver_sample.label or "")
solver_correct = _is_correct(solver_sample.response, ground_truth=ground_truth)
solver_sample.reward = 1.0 if solver_correct else 0.0
verifier_sample.reward = _compute_verifier_reward(
solver_correct=solver_correct,
verdict=_parse_verdict(verifier_sample.response),
verifier_correct=_is_correct(verifier_sample.response, ground_truth=ground_truth),
)
solver_sample.trainer_model_id = solver_model_id
verifier_sample.trainer_model_id = verifier_model_id
return GenerateFnOutput(samples=[solver_sample, verifier_sample])
def _compute_verifier_reward(*, solver_correct: bool, verdict: _Verdict | None, verifier_correct: bool) -> float:
if verdict is None:
return 0.0
if solver_correct:
return 1.0 if verdict is _Verdict.AGREE else 0.0
else:
if verdict is _Verdict.AGREE:
return 0.0
return 1.0 if verifier_correct else 0.5
def _parse_verdict(response: str) -> _Verdict | None:
if len(found := _VERDICT_PATTERN.findall(response)) != 1:
return None
return _Verdict(found[0])
def _compute_router_url(args, *, model_id: str) -> str:
host, port = args.sglang_model_routers[model_id]
return f"http://{host}:{port}/generate"
def _build_verifier_sample(solver_sample: Sample) -> Sample:
prompt = _VERIFIER_PROMPT_TEMPLATE.format(
question=_extract_question(solver_sample.prompt), solver_response=solver_sample.response
)
return Sample(
group_index=solver_sample.group_index,
index=solver_sample.index,
rollout_id=solver_sample.rollout_id,
prompt=[dict(role="user", content=prompt)],
label=solver_sample.label,
metadata=dict(solver_sample.metadata or {}),
routing_key=solver_sample.routing_key,
)
def _extract_question(prompt: str | list[dict[str, str]]) -> str:
assert not isinstance(prompt, str), (
"examples/multi_policy/solver_verifier.py quotes the question inside the verifier prompt, and a raw "
"string prompt may already be chat templated with special tokens, so the dataset must use messages"
)
[user_content] = [message["content"] for message in prompt if message["role"] == "user"]
return user_content
def _is_correct(response: str, *, ground_truth: str | None) -> bool:
if ground_truth is None:
return False
return _extract_answer(response) == ground_truth
def _extract_answer(text: str) -> str | None:
if found := _MARKED_ANSWER_PATTERN.findall(text):
return _normalize_answer(found[-1])
if found := _NUMBER_PATTERN.findall(text):
return _normalize_answer(found[-1])
return None
def _normalize_answer(answer: str) -> str:
return answer.strip().rstrip(".").replace(",", "").replace("$", "").strip()
+1
View File
@@ -33,6 +33,7 @@ KNOWN_LABELS: dict[str, str] = {
"ft-long": "Fault-tolerance trainer soak tests (random-crash survival, realistic-gsm8k convergence)",
"weight-update": "Weight update tests",
"fully-async": "Fully-async rollout tests",
"multi-policy": "Multi policy training tests (several policy models in one run)",
"replay": "Routing / indexer replay tests",
"qwen35": "Qwen3.5-35B-A3B MTP / spec-v2 e2e tests",
"mooncake": "Mooncake object-store rollout transfer tests",
@@ -0,0 +1,27 @@
import os
from tests.ci.ci_register import register_cuda_ci
from tests.e2e.short.test_multi_policy_solver_verifier_gsm8k import (
SOLVER_MODEL_ID,
VERIFIER_MODEL_ID,
TrainRewardBounds,
execute,
prepare,
)
register_cuda_ci(est_time=5400, suite="stage-c-8-gpu-h100", labels=["long"])
NUM_ROLLOUT = int(os.environ.get("MILES_TEST_NUM_ROLLOUT", "100"))
# TODO: tighten these weak bounds once the e2e run has been observed.
TRAIN_REWARD_BOUNDS = {
SOLVER_MODEL_ID: TrainRewardBounds(initial_max=0.6, final_min=0.5),
VERIFIER_MODEL_ID: TrainRewardBounds(initial_max=0.9, final_min=0.1),
}
if __name__ == "__main__":
prepare()
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
os.environ.pop(proxy_var, None)
execute(num_rollout=NUM_ROLLOUT, train_reward_bounds=TRAIN_REWARD_BOUNDS)
@@ -0,0 +1,233 @@
import os
import statistics
from pathlib import Path
from typing import NamedTuple
import yaml
from tests.ci.ci_register import register_cuda_ci
from miles.utils.audit_utils.event_logger.logger import read_events
from miles.utils.audit_utils.event_logger.models import MetricEvent
from miles.utils.external_utils import command_utils
from miles.utils.external_utils.command_utils.base_backend import ExecuteTrainConfig
from miles.utils.external_utils.command_utils.common import compute_model_args_overrides, encode_pseudo_file
register_cuda_ci(est_time=900, suite="stage-c-8-gpu-h100", labels=["short", "multi-policy", "fully-async"])
SOLVER_MODEL_ID = "solver"
VERIFIER_MODEL_ID = "verifier"
SOLVER_MODEL_NAME = "Qwen2.5-0.5B-Instruct"
VERIFIER_MODEL_NAME = "Qwen3-0.6B"
SOLVER_MODEL_TYPE = "qwen2.5-0.5B"
VERIFIER_MODEL_TYPE = "qwen3-0.6B"
NUM_GPUS = 8
NUM_ROLLOUT = int(os.environ.get("MILES_TEST_NUM_ROLLOUT", "3"))
SOLVER_PATH = f"/root/models/{SOLVER_MODEL_NAME}"
VERIFIER_PATH = f"/root/models/{VERIFIER_MODEL_NAME}"
# Hyperparameters roughly follow test_qwen2.5_0.5B_gsm8k_async.py.
SHARED_TRAINER_OVERRIDES = dict(
lr_decay_style="constant",
weight_decay=0.1,
adam_beta1=0.9,
adam_beta2=0.98,
kl_loss_coef=0.0,
kl_loss_type="low_var_kl",
entropy_coef=0.0,
eps_clip=0.2,
eps_clip_high=0.28,
)
SGLANG_CONFIG = dict(
sglang=[
dict(
name=SOLVER_MODEL_ID,
model_path=SOLVER_PATH,
update_weights=True,
num_gpus_per_engine=1,
server_groups=[dict(worker_type="regular", num_gpus=2)],
),
dict(
name=VERIFIER_MODEL_ID,
model_path=VERIFIER_PATH,
update_weights=True,
num_gpus_per_engine=1,
server_groups=[dict(worker_type="regular", num_gpus=2)],
),
]
)
class TrainRewardBounds(NamedTuple):
initial_max: float
final_min: float
TRAIN_REWARD_BOUNDS = {
SOLVER_MODEL_ID: TrainRewardBounds(initial_max=0.9, final_min=0.01),
VERIFIER_MODEL_ID: TrainRewardBounds(initial_max=0.9, final_min=0.01),
}
def prepare():
U = command_utils.default_config().create_backend()
U.exec_command_cpu("mkdir -p /root/models /root/datasets")
U.exec_command_cpu(f"hf download Qwen/{SOLVER_MODEL_NAME} --local-dir {SOLVER_PATH}")
U.exec_command_cpu(f"hf download Qwen/{VERIFIER_MODEL_NAME} --local-dir {VERIFIER_PATH}")
U.hf_download_dataset("zhuzilin/gsm8k")
def execute(*, num_rollout: int = NUM_ROLLOUT, train_reward_bounds: dict[str, TrainRewardBounds] | None = None):
config = command_utils.default_config()
U = config.create_backend()
events_dir = compute_events_dir(config)
megatron_config = compute_megatron_config()
ckpt_args = f"--hf-checkpoint {SOLVER_PATH}/ " f"--ref-load {SOLVER_PATH}/ "
policy_args = (
f"--megatron-config {encode_pseudo_file(yaml.dump(megatron_config))} "
f"--sglang-config {encode_pseudo_file(yaml.dump(SGLANG_CONFIG))} "
"--custom-generate-function-path examples.multi_policy.solver_verifier.generate "
)
rollout_args = (
"--fully-async "
"--prompt-data /root/datasets/gsm8k/train.parquet "
"--input-key messages "
"--label-key label "
"--apply-chat-template "
"--rollout-shuffle "
f"--num-rollout {num_rollout} "
"--rollout-batch-size 8 "
"--n-samples-per-prompt 4 "
"--rollout-max-response-len 1024 "
"--rollout-temperature 0.8 "
"--global-batch-size 32 "
# retract (default) can deadlock flush_cache in fully_async under load
"--pause-generation-mode in_place "
)
perf_args = (
"--tensor-model-parallel-size 1 "
"--sequence-parallel "
"--pipeline-model-parallel-size 1 "
"--context-parallel-size 1 "
"--expert-model-parallel-size 1 "
"--expert-tensor-parallel-size 1 "
"--use-dynamic-batch-size "
"--max-tokens-per-gpu 9216 "
)
grpo_args = "--advantage-estimator grpo " "--use-kl-loss "
optimizer_args = "--optimizer adam " "--lr 1e-6 "
sglang_args = "--rollout-num-gpus-per-engine 1 " "--sglang-mem-fraction-static 0.65 " "--sglang-enable-metrics "
ci_args = "--ci-test " f"--save-debug-event-data {events_dir} "
misc_args = (
"--attention-dropout 0.0 "
"--hidden-dropout 0.0 "
"--accumulate-allreduce-grads-in-fp32 "
"--attention-softmax-in-fp32 "
"--attention-backend flash "
"--actor-num-nodes 1 "
"--actor-num-gpus-per-node 2 "
"--rollout-num-gpus 4 "
"--megatron-to-hf-mode bridge "
)
train_args = (
f"{ckpt_args} "
f"{policy_args} "
f"{rollout_args} "
f"{optimizer_args} "
f"{grpo_args} "
f"{command_utils.get_default_wandb_args(__file__)} "
f"{perf_args} "
f"{sglang_args} "
f"{ci_args} "
f"{misc_args} "
)
U.execute_train(
train_args=train_args,
num_gpus_per_node=NUM_GPUS,
megatron_model_type=SOLVER_MODEL_TYPE,
train_script="train_multi_policy.py",
extra_env_vars={"MILES_EXPERIMENTAL_ROLLOUT_REFACTOR": "1"},
)
_assert_every_policy_learned(events_dir, bounds=train_reward_bounds or TRAIN_REWARD_BOUNDS)
def compute_megatron_config() -> dict:
return dict(
trainers=[
_compute_trainer_config(model_id=SOLVER_MODEL_ID, model_type=SOLVER_MODEL_TYPE, model_path=SOLVER_PATH),
_compute_trainer_config(
model_id=VERIFIER_MODEL_ID, model_type=VERIFIER_MODEL_TYPE, model_path=VERIFIER_PATH
),
]
)
def _compute_trainer_config(*, model_id: str, model_type: str, model_path: str) -> dict:
return dict(
model_id=model_id,
overrides=dict(
hf_checkpoint=model_path,
ref_load=model_path,
**compute_model_args_overrides(model_type),
**SHARED_TRAINER_OVERRIDES,
),
)
def compute_events_dir(config: ExecuteTrainConfig) -> Path:
return Path(config.output_dir) / "multi_policy_solver_verifier" / config.run_id / "events"
def _assert_every_policy_learned(events_dir: Path, *, bounds: dict[str, TrainRewardBounds]) -> None:
for model_id, model_bounds in bounds.items():
rewards = _read_train_reward_series(events_dir, model_id=model_id)
assert rewards, (
f"no {_compute_train_reward_key(model_id)} value was logged under {events_dir}, so policy "
f"{model_id!r} never reported a training reward and nothing about its learning can be checked"
)
initial = rewards[0]
final = statistics.mean(rewards[-max(1, len(rewards) // 3) :])
assert initial <= model_bounds.initial_max, (
f"policy {model_id!r} starts at training reward {initial}, above {model_bounds.initial_max}; a run "
f"that starts already solved cannot show that training moved it"
)
assert final >= model_bounds.final_min, (
f"policy {model_id!r} ends at training reward {final}, below {model_bounds.final_min}; either its "
f"reward function never fires, or training destroyed the model"
)
def _read_train_reward_series(events_dir: Path, *, model_id: str) -> list[float]:
reward_key = _compute_train_reward_key(model_id)
step_key = f"{model_id}/rollout/step"
points = [
(event.metrics[step_key], event.metrics[reward_key])
for event in read_events(events_dir)
if isinstance(event, MetricEvent) and reward_key in event.metrics
]
return [reward for _, reward in sorted(points)]
def _compute_train_reward_key(model_id: str) -> str:
return f"{model_id}/rollout/raw_reward"
if __name__ == "__main__":
prepare()
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
os.environ.pop(proxy_var, None)
execute()
@@ -1,3 +1,5 @@
import argparse
import re
import sys
from argparse import Namespace
from types import SimpleNamespace
@@ -10,6 +12,7 @@ from tests.fast.fixtures.megatron_config_fixtures import encode_megatron_config
from miles.backends.megatron_utils import megatron_config as megatron_config_module
from miles.backends.megatron_utils.megatron_config import (
MODEL_DEFINITION_ARGS,
PER_POLICY_ARGS,
_compute_trainer_checkpoint_dir,
_has_megatron_checkpoint,
@@ -19,6 +22,7 @@ from miles.backends.megatron_utils.megatron_config import (
resolve_args_checkpoint_load,
resolve_megatron_config,
)
from miles.utils.external_utils.model_args_utils import load_model_args
from miles.utils.workers.naming import TRAINER_ID_MAX_LENGTH
@@ -271,6 +275,67 @@ class TestOverrideCoercion:
resolve_megatron_config(_make_args(path))
class TestModelDefinitionOverrides:
@staticmethod
def _parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(add_help=False)
parser.add_argument("--lr", type=float)
parser.add_argument("--num-layers", type=int)
parser.add_argument("--num-experts", type=int)
parser.add_argument("--moe-router-topk", type=int)
parser.add_argument("--untie-embeddings-and-output-weights", action="store_true")
parser.add_argument("--spec", nargs="*")
parser.add_argument("--num-rollout", type=int)
return parser
@pytest.fixture(autouse=True)
def _parser_of_run(self, monkeypatch):
monkeypatch.setattr(megatron_config_module, "get_megatron_arg_parser", self._parser)
def test_a_model_definition_argument_this_run_declares_is_accepted_and_typed(self):
"""A policy carries the shape of its own model, of which PER_POLICY_ARGS lists only a part."""
overrides = _resolve_overrides({"num_experts": "8", "moe_router_topk": "4"}, model_id="a")
assert overrides == {"num_experts": 8, "moe_router_topk": 4}
def test_a_model_definition_flag_stays_a_boolean(self):
"""store_true arguments carry no value on a command line, so the overlay must keep the boolean."""
overrides = _resolve_overrides({"untie_embeddings_and_output_weights": True}, model_id="a")
assert overrides == {"untie_embeddings_and_output_weights": True}
def test_a_model_definition_argument_this_run_does_not_declare_is_refused(self):
"""The set is intersected with the parser, so a name this run cannot type admits nothing."""
with pytest.raises(AssertionError, match="may not override"):
_resolve_overrides({"kv_lora_rank": 512}, model_id="a")
def test_an_argument_in_neither_whitelist_is_refused(self):
"""The parser declares --num-rollout, yet the rhythm of a run is read from the base command line."""
with pytest.raises(AssertionError, match="num_rollout"):
_resolve_overrides({"num_rollout": 3}, model_id="a")
def test_a_list_reaches_an_argument_that_takes_several_values(self):
"""--spec names a module and a function, which is how such an argument arrives from yaml."""
overrides = _resolve_overrides({"spec": ["miles_plugins.models.glm5.glm5", "get_glm5_spec"]}, model_id="a")
assert overrides == {"spec": ["miles_plugins.models.glm5.glm5", "get_glm5_spec"]}
def test_a_list_written_where_a_single_value_is_declared_is_refused(self):
"""One integer argument cannot be given two, and the overlay is the only place that can say so."""
with pytest.raises(AssertionError, match="takes a single value"):
_resolve_overrides({"num_experts": [8, 4]}, model_id="a")
def test_a_training_knob_is_still_accepted_beside_a_model_definition_argument(self):
"""The model definition arguments are admitted on top of PER_POLICY_ARGS, not instead of it."""
overrides = _resolve_overrides({"lr": "5e-7", "num_layers": "12"}, model_id="a")
assert overrides == {"lr": 5e-7, "num_layers": 12}
def test_the_deprecated_parser_epsilon_destination_is_whitelisted(self) -> None:
"""The deprecated parser stores its epsilon option under norm_epsilon, which policies may override."""
assert "norm_epsilon" in MODEL_DEFINITION_ARGS
class TestResolveOverrides:
def test_an_empty_override_map_never_builds_the_parser(self, monkeypatch):
"""Building the parser imports and runs megatron's whole argument stack, per trainer that overrides nothing."""
@@ -781,3 +846,47 @@ class TestBaseArgumentsAreNotMutated:
model.train_env_vars["NCCL_DEBUG"] = "INFO"
assert args.train_env_vars == {"NCCL_DEBUG": "WARN"}
class TestPerPolicyArgsCoverage:
@pytest.mark.parametrize("model_type", ["qwen2.5-0.5B", "qwen3-0.6B"])
def test_the_whitelist_admits_every_argument_of_a_model_script(self, model_type):
"""A policy that cannot override one of its own architecture arguments would train another model's shape."""
options = get_megatron_arg_parser()._option_string_actions
flags = [token for token in load_model_args(model_type).split() if token.startswith("--")]
unknown = [flag for flag in flags if flag not in options]
assert not unknown, f"{model_type} passes {unknown}, which the megatron parser does not declare"
assert {options[flag].dest for flag in flags} <= PER_POLICY_ARGS | MODEL_DEFINITION_ARGS
@pytest.mark.parametrize("model_type", ["qwen2.5-0.5B", "qwen3-0.6B"])
def test_a_model_script_names_only_model_definition_arguments(self, model_type):
"""PER_POLICY_ARGS holds the training knobs, so nothing describing a model shape may hide in it."""
options = get_megatron_arg_parser()._option_string_actions
flags = [token for token in load_model_args(model_type).split() if token.startswith("--")]
assert {options[flag].dest for flag in flags} <= MODEL_DEFINITION_ARGS
def test_the_whitelist_names_arguments_the_parser_actually_produces(self):
"""The override keys are compared against this set and then looked up by the same name in the
parser, so a flag spelled the way it appears on a command line admits nothing at all: a config
carrying the argument's real name is refused, and the name in the set can never be reached."""
parsed = {action.dest for action in get_megatron_arg_parser()._actions}
assert PER_POLICY_ARGS <= parsed
class TestOverrideWhitelistShape:
def test_the_two_whitelists_are_disjoint(self):
"""A name in both no longer says whether it is a training knob or part of a model definition."""
assert not PER_POLICY_ARGS & MODEL_DEFINITION_ARGS
def test_every_whitelisted_name_is_spelled_the_way_a_destination_is(self):
"""Both sets are matched against argparse destinations, so a command line spelling admits nothing."""
misspelled = [
name
for name in sorted(PER_POLICY_ARGS | MODEL_DEFINITION_ARGS)
if re.fullmatch(r"[a-z][a-z0-9_]*", name) is None
]
assert misspelled == []
@@ -0,0 +1,354 @@
from argparse import Namespace
from dataclasses import dataclass, field
from pathlib import Path
from types import SimpleNamespace
import pytest
from examples.multi_policy import solver_verifier
from examples.multi_policy.solver_verifier import _Verdict
from tests.fast.fixtures.megatron_config_fixtures import encode_megatron_config
from miles.rollout.base_types import GenerateFnInput, GenerateFnOutput
from miles.utils.types import Sample
SOLVER_URL = "http://solver-host:1111/generate"
VERIFIER_URL = "http://verifier-host:2222/generate"
@dataclass
class _FakeGenerate:
responses: dict[str, str]
calls: list[tuple[str, Sample]] = field(default_factory=list)
async def __call__(self, input: GenerateFnInput, url: str | None = None) -> GenerateFnOutput:
sample = input.sample
self.calls.append((url, sample))
sample.response = self.responses[url]
sample.status = Sample.Status.COMPLETED
return GenerateFnOutput(samples=sample)
def _make_input(*, prompt: str | list[dict[str, str]], label: str) -> GenerateFnInput:
args = Namespace(
megatron_config=encode_megatron_config("solver", "verifier"),
use_critic=False,
sglang_model_routers={"solver": ("solver-host", 1111), "verifier": ("verifier-host", 2222)},
)
sample = Sample(group_index=3, index=7, prompt=prompt, label=label)
return GenerateFnInput(state=SimpleNamespace(args=args), sample=sample, sampling_params={}, evaluation=False)
@dataclass(frozen=True)
class _RunResult:
fake: _FakeGenerate
samples: list[Sample]
async def _run(monkeypatch, *, solver_response: str, verifier_response: str) -> _RunResult:
fake = _FakeGenerate(responses={SOLVER_URL: solver_response, VERIFIER_URL: verifier_response})
monkeypatch.setattr(solver_verifier, "single_turn_generate", fake)
output = await solver_verifier.generate(
_make_input(prompt=[dict(role="user", content="What is 9 + 9?")], label="#### 18")
)
return _RunResult(fake=fake, samples=output.samples)
class TestComputeVerifierReward:
@pytest.mark.parametrize("verifier_correct", [False, True])
def test_agreeing_with_a_right_solver_is_the_only_full_credit_case(self, verifier_correct):
"""The solver was right and the verifier said so, so its own answer never enters the score."""
assert (
solver_verifier._compute_verifier_reward(
solver_correct=True, verdict=_Verdict.AGREE, verifier_correct=verifier_correct
)
== 1.0
)
@pytest.mark.parametrize("verifier_correct", [False, True])
def test_calling_a_right_solver_wrong_scores_zero(self, verifier_correct):
"""A false accusation is worthless however good the verifier's replacement answer is."""
assert (
solver_verifier._compute_verifier_reward(
solver_correct=True, verdict=_Verdict.WRONG, verifier_correct=verifier_correct
)
== 0.0
)
@pytest.mark.parametrize("verifier_correct", [False, True])
def test_agreeing_with_a_wrong_solver_scores_zero(self, verifier_correct):
"""Endorsing a wrong solution is the failure the verifier exists to avoid."""
assert (
solver_verifier._compute_verifier_reward(
solver_correct=False, verdict=_Verdict.AGREE, verifier_correct=verifier_correct
)
== 0.0
)
def test_catching_a_wrong_solver_without_fixing_it_scores_half(self):
"""Spotting the error is worth partial credit even when the replacement answer is wrong."""
assert (
solver_verifier._compute_verifier_reward(
solver_correct=False, verdict=_Verdict.WRONG, verifier_correct=False
)
== 0.5
)
def test_catching_a_wrong_solver_and_fixing_it_scores_full(self):
"""Both halves of the verifier's job were done."""
assert (
solver_verifier._compute_verifier_reward(
solver_correct=False, verdict=_Verdict.WRONG, verifier_correct=True
)
== 1.0
)
@pytest.mark.parametrize("solver_correct", [False, True])
@pytest.mark.parametrize("verifier_correct", [False, True])
def test_an_unparseable_verdict_scores_zero(self, solver_correct, verifier_correct):
"""A verdict nobody can read teaches the solver nothing, whatever the verifier meant."""
assert (
solver_verifier._compute_verifier_reward(
solver_correct=solver_correct, verdict=None, verifier_correct=verifier_correct
)
== 0.0
)
class TestParseVerdict:
def test_the_verdict_line_the_prompt_asks_for_is_read(self):
"""The happy path the verifier prompt asks for."""
assert solver_verifier._parse_verdict("The arithmetic checks out.\nVERDICT: AGREE") is _Verdict.AGREE
def test_a_marker_inside_the_reasoning_is_not_a_verdict(self):
"""Otherwise 'I cannot decide: AGREE or WRONG' scores as if the verifier had ruled."""
assert solver_verifier._parse_verdict("I would AGREE, but the sum is off.\nWRONG\n#### 18") is None
def test_two_verdict_lines_are_unparseable(self):
"""A reply that rules twice never made up its mind, and picking one of them invents a decision."""
response = "VERDICT: AGREE\nOn reflection:\nVERDICT: WRONG"
assert solver_verifier._parse_verdict(response) is None
def test_a_lowercase_verdict_is_not_a_verdict(self):
"""Strict markers keep prose such as 'I agree with the setup' from scoring."""
assert solver_verifier._parse_verdict("verdict: agree") is None
def test_a_longer_word_containing_the_marker_is_not_a_verdict(self):
"""The verdict line holds the marker alone, so 'AGREEMENT' cannot be read as one."""
assert solver_verifier._parse_verdict("VERDICT: AGREEMENT") is None
def test_a_response_without_any_marker_is_unparseable(self):
"""An empty or rambling reply has no verdict to score."""
assert solver_verifier._parse_verdict("") is None
def test_the_wrong_verdict_line_is_read(self):
"""The other half of the protocol, and the only path that lets the verifier earn credit on a bad solution."""
assert solver_verifier._parse_verdict("The sum is off.\nVERDICT: WRONG\n#### 18") is _Verdict.WRONG
@pytest.mark.parametrize("response", ["VERDICT:AGREE", "VERDICT: AGREE "])
def test_spacing_around_the_marker_is_tolerated(self, response):
"""Models are inconsistent about the space after the colon, and that is not a decision."""
assert solver_verifier._parse_verdict(response) is _Verdict.AGREE
def test_an_indented_verdict_line_is_not_a_verdict(self):
"""The prompt asks for a line of its own, so an indented one is quoted text rather than a ruling."""
assert solver_verifier._parse_verdict("Example:\n VERDICT: AGREE") is None
def test_trailing_text_on_the_verdict_line_is_not_a_verdict(self):
"""'VERDICT: AGREE with reservations' is a sentence, and reading it as a ruling invents certainty."""
assert solver_verifier._parse_verdict("VERDICT: AGREE with reservations") is None
def test_the_same_verdict_twice_is_still_unparseable(self):
"""Counting lines, not distinct values, keeps a repeated ruling from looking more decided than it is."""
assert solver_verifier._parse_verdict("VERDICT: AGREE\nVERDICT: AGREE") is None
class TestExtractAnswer:
def test_the_gsm8k_marker_is_read(self):
"""Both the dataset label and the prompted reply end with '#### <answer>'."""
assert solver_verifier._extract_answer("Half of 36 is 18.\n#### 18") == "18"
def test_the_last_marker_wins(self):
"""A reply that reconsiders itself is scored on its final answer."""
assert solver_verifier._extract_answer("#### 17\nOn reflection:\n#### 18") == "18"
def test_a_marked_answer_is_normalized(self):
"""Currency, thousand separators and a trailing period are formatting, not the answer."""
assert solver_verifier._extract_answer("#### $1,234.") == "1234"
def test_a_reply_without_the_marker_falls_back_to_its_last_number(self):
"""The solver prompt comes from the dataset, so it need not ask for the marker."""
assert solver_verifier._extract_answer("First 9, then 9, so the total is 18") == "18"
def test_a_reply_without_any_number_has_no_answer(self):
"""Nothing to compare against the ground truth."""
assert solver_verifier._extract_answer("I cannot tell") is None
def test_a_negative_marked_answer_keeps_its_sign(self):
"""gsm8k answers can be negative, and dropping the sign would score a wrong reply as right."""
assert solver_verifier._extract_answer("#### -5") == "-5"
def test_the_fallback_reads_the_last_number_not_the_first(self):
"""An unmarked reply states its result last, after restating the operands."""
assert solver_verifier._extract_answer("From 20 we subtract 2, so 18") == "18"
def test_the_marker_wins_over_a_later_bare_number(self):
"""Trailing prose such as a page reference must not overwrite the answer the model marked."""
assert solver_verifier._extract_answer("#### 18\nsee step 3") == "18"
class TestIsCorrect:
def test_a_label_without_an_answer_never_matches(self):
"""A dataset row with no parseable answer must fail closed rather than reward everything."""
assert not solver_verifier._is_correct("#### 18", ground_truth=None)
def test_both_sides_are_normalized_before_comparing(self):
"""The label and the reply format the same number differently, and formatting is not a wrong answer."""
assert solver_verifier._is_correct("#### $1,234.", ground_truth="1234")
class TestBuildVerifierSample:
def test_the_verifier_sample_keeps_the_identity_of_the_solver_sample(self):
"""Both samples belong to one trajectory, and the group is what advantage is computed over."""
solver = Sample(group_index=3, index=7, rollout_id=2, prompt=[dict(role="user", content="q")], label="#### 18")
solver.response = "#### 18"
verifier = solver_verifier._build_verifier_sample(solver)
assert (verifier.group_index, verifier.index, verifier.rollout_id) == (3, 7, 2)
assert verifier.label == "#### 18"
def test_the_metadata_is_copied_rather_than_shared(self):
"""The two samples are scored and dumped separately, so one must not write into the other's metadata."""
solver = Sample(prompt=[dict(role="user", content="q")], metadata={"source": "gsm8k"})
solver.response = "#### 18"
verifier = solver_verifier._build_verifier_sample(solver)
verifier.metadata["source"] = "mutated"
assert solver.metadata == {"source": "gsm8k"}
def test_the_routing_key_is_carried_over(self):
"""Consistent hashing routes both samples of a trajectory to the same engine."""
solver = Sample(prompt=[dict(role="user", content="q")], routing_key="key-1")
solver.response = "#### 18"
assert solver_verifier._build_verifier_sample(solver).routing_key == "key-1"
class TestExtractQuestion:
def test_a_system_message_is_not_quoted_as_the_question(self):
"""Only the user turn holds the problem; quoting the system prompt would confuse the verifier."""
prompt = [dict(role="system", content="Be brief."), dict(role="user", content="What is 9 + 9?")]
assert solver_verifier._extract_question(prompt) == "What is 9 + 9?"
def test_several_user_messages_are_refused(self):
"""This example assumes one question per sample, and silently picking one would mis-state the task."""
prompt = [dict(role="user", content="a"), dict(role="user", content="b")]
with pytest.raises(ValueError):
solver_verifier._extract_question(prompt)
class TestComputeRouterUrl:
def test_the_url_points_at_the_router_of_that_policy(self):
"""Nothing else routes by model id, so a wrong url would silently train one policy on the other's tokens."""
args = Namespace(sglang_model_routers={"solver": ("solver-host", 1111)})
assert solver_verifier._compute_router_url(args, model_id="solver") == SOLVER_URL
class TestGenerate:
async def test_the_verifier_prompt_quotes_the_question_and_the_solver_answer(self, monkeypatch):
"""The verifier only sees the solver's work through the prompt this function assembles."""
result = await _run(monkeypatch, solver_response="It is 18.\n#### 18", verifier_response="VERDICT: AGREE")
verifier_prompt = result.fake.calls[1][1].prompt
assert isinstance(verifier_prompt, list)
assert verifier_prompt[0]["role"] == "user"
assert "What is 9 + 9?" in verifier_prompt[0]["content"]
assert "It is 18.\n#### 18" in verifier_prompt[0]["content"]
async def test_a_raw_string_prompt_is_refused(self, monkeypatch):
"""A string prompt may already be chat templated, so quoting it as the question would leak tokens."""
fake = _FakeGenerate(responses={SOLVER_URL: "#### 18", VERIFIER_URL: "VERDICT: AGREE"})
monkeypatch.setattr(solver_verifier, "single_turn_generate", fake)
with pytest.raises(AssertionError, match="chat templated"):
await solver_verifier.generate(_make_input(prompt="What is 9 + 9?", label="#### 18"))
async def test_each_policy_is_generated_against_its_own_router(self, monkeypatch):
"""Nothing routes by trainer_model_id, so the generate function picks the url itself."""
result = await _run(monkeypatch, solver_response="#### 18", verifier_response="VERDICT: AGREE")
assert [url for url, _ in result.fake.calls] == [SOLVER_URL, VERIFIER_URL]
async def test_both_samples_are_returned_bound_to_their_own_policy(self, monkeypatch):
"""trainer_model_id is filled on return, and it is what sends each sample to its trainer."""
result = await _run(monkeypatch, solver_response="#### 18", verifier_response="VERDICT: AGREE")
solver_sample, verifier_sample = result.samples
assert solver_sample.trainer_model_id == "solver"
assert verifier_sample.trainer_model_id == "verifier"
async def test_a_right_solver_endorsed_by_the_verifier_rewards_both(self, monkeypatch):
"""The end to end path of the full credit row of the reward matrix."""
result = await _run(monkeypatch, solver_response="#### 18", verifier_response="Checks out.\nVERDICT: AGREE")
solver_sample, verifier_sample = result.samples
assert solver_sample.reward == 1.0
assert verifier_sample.reward == 1.0
async def test_a_wrong_solver_corrected_by_the_verifier_rewards_only_the_verifier(self, monkeypatch):
"""The solver is scored against the label, the verifier against what it did about the solver."""
result = await _run(monkeypatch, solver_response="#### 17", verifier_response="VERDICT: WRONG\n#### 18")
solver_sample, verifier_sample = result.samples
assert solver_sample.reward == 0.0
assert verifier_sample.reward == 1.0
async def test_a_wrong_solver_caught_but_not_fixed_rewards_half(self, monkeypatch):
"""The verifier's own answer is graded against the same ground truth."""
result = await _run(monkeypatch, solver_response="#### 17", verifier_response="VERDICT: WRONG\n#### 16")
assert result.samples[1].reward == 0.5
async def test_a_run_naming_one_policy_is_refused(self, monkeypatch):
"""This example needs a solver and a verifier, and it must not silently train one of them twice."""
fake = _FakeGenerate(responses={})
monkeypatch.setattr(solver_verifier, "single_turn_generate", fake)
input = _make_input(prompt=[dict(role="user", content="What is 9 + 9?")], label="#### 18")
input.args.megatron_config = encode_megatron_config("solver")
with pytest.raises(AssertionError, match="pairs one solver policy with one verifier policy"):
await solver_verifier.generate(input)
class TestTheLauncherLeavesThePromptAsMessages:
def test_the_e2e_test_does_not_ask_the_dataset_to_apply_the_chat_template(self):
"""--apply-chat-template renders the messages into one templated string at dataset build time.
This example quotes the question into a second prompt, which a string carrying special tokens
cannot be used for, and a list prompt is chat templated at generation anyway. Getting this wrong
costs an 8-GPU run to notice."""
source = (
Path(__file__).resolve().parents[4] / "examples/multi_policy/run_solver_verifier_gsm8k.py"
).read_text()
# the comment saying why the flag is absent names the flag, so read only what is handed to the run
launched = "\n".join(line for line in source.splitlines() if not line.strip().startswith("#"))
assert "--custom-generate-function-path examples.multi_policy.solver_verifier.generate" in launched
assert "--apply-chat-template" not in launched
def test_a_templated_string_prompt_is_refused_rather_than_quoted(self):
"""The refusal is what keeps a prompt full of special tokens out of the verifier's question."""
with pytest.raises(AssertionError, match="the dataset must use messages"):
solver_verifier._extract_question("<|im_start|>user\nWhat is 9 + 9?<|im_end|>\n")
def test_a_message_prompt_gives_up_its_question(self):
"""The other half: the shape the launcher now preserves is the one this reads."""
question = solver_verifier._extract_question(
[dict(role="system", content="be terse"), dict(role="user", content="What is 9 + 9?")]
)
assert question == "What is 9 + 9?"