mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Add a solver-verifier example on gsm8k (#2645)
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
@@ -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?"
|
||||
Reference in New Issue
Block a user