diff --git a/torchtitan/experiments/__init__.py b/torchtitan/experiments/__init__.py index 2e9918c7f1d..c430ba208f4 100644 --- a/torchtitan/experiments/__init__.py +++ b/torchtitan/experiments/__init__.py @@ -17,6 +17,7 @@ # RL examples own a per-example config_registry under rl/examples/; # listed here so `--module ` resolves (see ConfigManager). "alphabet_sort", + "dapo_math", "search_r1", ] ) diff --git a/torchtitan/experiments/rl/examples/dapo_math/README.md b/torchtitan/experiments/rl/examples/dapo_math/README.md new file mode 100644 index 00000000000..891243e7663 --- /dev/null +++ b/torchtitan/experiments/rl/examples/dapo_math/README.md @@ -0,0 +1,102 @@ +# DAPO Math + +[DAPO-Math-17k](https://huggingface.co/datasets/BytedTsinghua-SIA/DAPO-Math-17k) is the verifiable math dataset released with [DAPO](https://arxiv.org/abs/2503.14476). This environment trains Qwen3-4B-Base with DAPO loss on a filtered version of that dataset. + +## Environment + +Each episode is single-turn: + +```text +user math problem -> one assistant solution -> binary Math-Verify reward +``` + +The prompt asks for step-by-step reasoning followed by a final `Answer:` expression. [Math-Verify](https://github.com/huggingface/Math-Verify) parses that expression and assigns a reward of one when it is mathematically equivalent to the reference answer, or zero otherwise. + +An episode from the reference run is shown below. The prompt is reproduced in full; the response is abridged. + +```text +Prompt: +Solve the following math problem step by step. The last line of your response +should be of the form Answer: $Answer (without quotes) where $Answer is the +answer to the problem. + +Let $r_1, r_2, \ldots, r_{47}$ be the roots of $x^{47} - 1 = 0$. Compute +\( \sum_{i=1}^{47} r_i^{2020} \). + +Remember to put your answer on its own line after "Answer:". + +Response: +The roots are the 47th roots of unity. Since 2020 is congruent to -1 +modulo 47, raising every root to the 2020th power permutes the roots. +Their sum is therefore zero. + +Answer: \boxed{0} + +Reward: 1 +``` + +## Datasets + +Training uses the 12,643-row [filtered DAPO-Math dataset](https://huggingface.co/datasets/hamishivi/DAPO-Math-17k-Processed_filtered). Each row contains one user prompt and its verifiable final answer. + +Validation uses all 30 problems from [AIME 2025](https://huggingface.co/datasets/opencompass/AIME2025). The same single-turn environment and Math-Verify reward are used for training and validation. + +## Reference configurations + +Both configurations run 150 optimizer steps on one eight-GPU node. One TP=2 trainer uses two GPUs, and six independent TP=1 generators use the remaining GPUs. Each optimizer step consumes 8 prompt groups with 16 completions per group. `max_offpolicy_steps=4` bounds policy lag. + +The 8K configuration is the default reference recipe: + +```text +config: rl_dapo_qwen3_4b_math_8k +prompt budget: 2,048 tokens +response budget: 8,192 tokens +packing length: 10,240 tokens +``` + +The 32K configuration extends the response and packing budgets while keeping the same model, optimizer, and GPU topology: + +```text +config: rl_dapo_qwen3_4b_math_32k +prompt budget: 2,048 tokens +response budget: 32,768 tokens +packing length: 34,816 tokens +``` + +Both configurations use a constant learning rate of `1e-6`, DAPO clipping of `[0.2, 0.28]`, and an fp32 language-model head with a bf16 model forward. The 32K configuration has not been benchmarked. + +## Setup + +Follow the [RL environment setup](../../README.md), install this recipe's verifier, and download the base checkpoint: + +```bash +pip install -r torchtitan/experiments/rl/examples/dapo_math/requirements.txt + +python scripts/download_hf_assets.py \ + --repo_id Qwen/Qwen3-4B-Base \ + --local_dir torchtitan/experiments/rl/example_checkpoint \ + --all +``` + +## Run + +Run the 150-step 8K reference configuration from the repository root. CLI arguments override fields from the config registry; this example selects an explicit output directory: + +```bash +python -m torchtitan.experiments.rl.train \ + --module dapo_math \ + --config rl_dapo_qwen3_4b_math_8k \ + --dump-folder outputs/rl/qwen3_4b_dapo_math_8k_150 +``` + +Use `rl_dapo_qwen3_4b_math_32k` as the config name to run the 32K variant. + +## 150-step reference result + +TODO: add eval results. Add 32k variant. + +The plots below were generated on July 20, 2026 from commit [`fec3e196`](https://github.com/felipemello1/torchtitan/commit/fec3e196a4ceb87bfc87fb4f1a36a538d7e98ee4). + +![Qwen3-4B DAPO Math-Verify reward](./assets/qwen3_4b_7k_reward.png) + +![Qwen3-4B DAPO mean response length](./assets/qwen3_4b_7k_response_length.png) diff --git a/torchtitan/experiments/rl/examples/dapo_math/__init__.py b/torchtitan/experiments/rl/examples/dapo_math/__init__.py new file mode 100644 index 00000000000..b337a213f9e --- /dev/null +++ b/torchtitan/experiments/rl/examples/dapo_math/__init__.py @@ -0,0 +1,27 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from torchtitan.experiments.rl.examples.dapo_math.data import ( + AIME2025Dataset, + DapoMathDataset, + DapoMathSample, +) +from torchtitan.experiments.rl.examples.dapo_math.env import DapoMathEnv +from torchtitan.experiments.rl.examples.dapo_math.rollouter import DapoMathRollouter +from torchtitan.experiments.rl.examples.dapo_math.rubric import ( + RewardMathVerify, + score_math_response, +) + +__all__ = [ + "AIME2025Dataset", + "DapoMathDataset", + "DapoMathEnv", + "DapoMathRollouter", + "DapoMathSample", + "RewardMathVerify", + "score_math_response", +] diff --git a/torchtitan/experiments/rl/examples/dapo_math/assets/qwen3_4b_7k_response_length.png b/torchtitan/experiments/rl/examples/dapo_math/assets/qwen3_4b_7k_response_length.png new file mode 100644 index 00000000000..423acf12542 Binary files /dev/null and b/torchtitan/experiments/rl/examples/dapo_math/assets/qwen3_4b_7k_response_length.png differ diff --git a/torchtitan/experiments/rl/examples/dapo_math/assets/qwen3_4b_7k_reward.png b/torchtitan/experiments/rl/examples/dapo_math/assets/qwen3_4b_7k_reward.png new file mode 100644 index 00000000000..f79371fc2f6 Binary files /dev/null and b/torchtitan/experiments/rl/examples/dapo_math/assets/qwen3_4b_7k_reward.png differ diff --git a/torchtitan/experiments/rl/examples/dapo_math/config_registry.py b/torchtitan/experiments/rl/examples/dapo_math/config_registry.py new file mode 100644 index 00000000000..0199adf5a1a --- /dev/null +++ b/torchtitan/experiments/rl/examples/dapo_math/config_registry.py @@ -0,0 +1,161 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Single-node Qwen3-4B-Base DAPO-Math recipes.""" + +from __future__ import annotations + +from torchtitan.components.checkpoint import CheckpointManager +from torchtitan.components.loss import ChunkedLossWrapper +from torchtitan.components.lr_scheduler import LRSchedulersContainer +from torchtitan.components.optimizer import default_adamw +from torchtitan.config import CompileConfig, ParallelismConfig, TrainingConfig +from torchtitan.experiments.rl.actors.generator import ( + SamplingConfig, + VLLMCudagraphConfig, + VLLMGenerator, +) +from torchtitan.experiments.rl.actors.trainer import PolicyTrainer +from torchtitan.experiments.rl.components.batcher import BatchConfig, Batcher +from torchtitan.experiments.rl.controller import ( + AsyncLoopConfig, + Controller, + ValidationConfig, +) +from torchtitan.experiments.rl.environment import TokenEnv +from torchtitan.experiments.rl.examples.dapo_math.data import AIME2025Dataset +from torchtitan.experiments.rl.examples.dapo_math.rollouter import DapoMathRollouter +from torchtitan.experiments.rl.losses import DAPOLoss +from torchtitan.experiments.rl.models.cast_linear import LMHeadCastConverter +from torchtitan.experiments.rl.models.vllm_registry import InferenceParallelismConfig +from torchtitan.experiments.rl.observability.metrics import MetricsProcessor +from torchtitan.experiments.rl.renderer import RendererConfig +from torchtitan.experiments.rl.routing.inter_generator_router import ( + InterGeneratorRouter, +) +from torchtitan.experiments.rl.routing.strategies import LeastLoadedRoutingStrategy +from torchtitan.models.qwen3 import model_registry + + +def _qwen3_4b_dapo_math_config( + *, + max_response_tokens: int, + max_total_tokens: int, + dump_folder: str, +) -> Controller.Config: + """Build the shared Qwen3-4B DAPO-Math configuration.""" + num_validation_samples = 30 + validation_dataset = AIME2025Dataset.Config( + num_samples=num_validation_samples, + ) + return Controller.Config( + model_spec=model_registry( + "4B", + attn_backend="varlen", + # Compute vocabulary logits in fp32; the rest of the forward uses bf16. + converters=[LMHeadCastConverter.Config()], + ), + hf_assets_path="torchtitan/experiments/rl/example_checkpoint/Qwen3-4B-Base", + dump_folder=dump_folder, + async_loop=AsyncLoopConfig( + num_training_steps=150, + num_groups_per_train_step=8, + group_size=16, + max_offpolicy_steps=4, + validation=ValidationConfig( + num_samples=num_validation_samples, + ), + batcher=Batcher.Config( + batch=BatchConfig(local_batch_size=1, seq_len=max_total_tokens), + ), + ), + compile=CompileConfig(enable=True, backend="aot_eager"), + rollouter=DapoMathRollouter.Config( + validation_dataset=validation_dataset, + token_env=TokenEnv.Config( + max_rollout_tokens=max_total_tokens, + max_num_turns=1, + ), + ), + renderer=RendererConfig(name="qwen3", enable_thinking=True), + num_generators=6, + generator_router=InterGeneratorRouter.Config( + strategy=LeastLoadedRoutingStrategy.Config() + ), + metrics=MetricsProcessor.Config( + enable_wandb=True, + console_log_keys_validation=[ + "validation_reward/_mean", + "validation_reward/_max", + "validation/response_length/mean", + "timing/validate", + ], + ), + trainer=PolicyTrainer.Config( + optimizer=default_adamw( + lr=1e-6, + betas=(0.9, 0.98), + weight_decay=0.1, + ), + # A minimum factor of 1 keeps the learning rate constant. + lr_scheduler=LRSchedulersContainer.Config( + warmup_steps=0, + min_lr_factor=1.0, + ), + training=TrainingConfig(), + parallelism=ParallelismConfig( + data_parallel_replicate_degree=1, + data_parallel_shard_degree=1, + tensor_parallel_degree=2, + ), + checkpoint=CheckpointManager.Config( + enable=True, + initial_load_in_hf=True, + interval=100, + last_save_model_only=False, + keep_latest_k=3, + ), + loss=ChunkedLossWrapper.Config( + num_chunks=8, + loss_fn=DAPOLoss.Config( + ratio_clip_low=0.2, + ratio_clip_high=0.28, + ), + ), + ), + generator=VLLMGenerator.Config( + model_dtype="bfloat16", + parallelism=InferenceParallelismConfig( + data_parallel_degree=1, + tensor_parallel_degree=1, + ), + cudagraph=VLLMCudagraphConfig(enable=True), + checkpoint=CheckpointManager.Config(enable=False), + sampling=SamplingConfig( + temperature=1.0, + top_p=1.0, + max_tokens=max_response_tokens, + ), + ), + ) + + +def rl_dapo_qwen3_4b_math_8k() -> Controller.Config: + """Run 8K responses on one node: one TP=2 trainer and six TP=1 generators.""" + return _qwen3_4b_dapo_math_config( + max_response_tokens=8192, + max_total_tokens=10240, + dump_folder="outputs/rl/qwen3_4b_dapo_math_8k", + ) + + +def rl_dapo_qwen3_4b_math_32k() -> Controller.Config: + """Run 32K responses on one node: one TP=2 trainer and six TP=1 generators.""" + return _qwen3_4b_dapo_math_config( + max_response_tokens=32768, + max_total_tokens=34816, + dump_folder="outputs/rl/qwen3_4b_dapo_math_32k", + ) diff --git a/torchtitan/experiments/rl/examples/dapo_math/data.py b/torchtitan/experiments/rl/examples/dapo_math/data.py new file mode 100644 index 00000000000..5146e908aa3 --- /dev/null +++ b/torchtitan/experiments/rl/examples/dapo_math/data.py @@ -0,0 +1,138 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +import random +from collections.abc import Iterator +from dataclasses import dataclass + +from datasets import concatenate_datasets, load_dataset + +from torchtitan.config import Configurable + + +_AIME_PROMPT_TEMPLATE = ( + "Solve the following math problem step by step. The last line of your response " + "should be of the form Answer: $Answer (without quotes) where $Answer is the " + "answer to the problem.\n\n" + "{problem}\n\n" + 'Remember to put your answer on its own line after "Answer:".' +) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class DapoMathSample: + """A math prompt paired with its expected final answer.""" + + prompt: str + ground_truth: str + + +# TODO: Share this cycling iterator with other RL datasets instead of keeping +# per-environment implementations. +class _CyclingDataset(Configurable): + """Provides an endless, resumable stream over a finite sample list.""" + + def __init__( + self, + samples: list[DapoMathSample], + *, + seed: int, + shuffle: bool, + ) -> None: + if not samples: + raise ValueError("math dataset must contain at least one sample") + self._samples = samples + self._rng = random.Random(seed) + self._shuffle = shuffle + self._order = list(range(len(samples))) + if shuffle: + self._rng.shuffle(self._order) + self._position = 0 + + def __iter__(self) -> Iterator[DapoMathSample]: + return self + + def __next__(self) -> DapoMathSample: + if self._position == len(self._order): + # Rollout production consumes an endless stream; crossing the dataset + # boundary starts a new epoch. Training reshuffles; validation does not. + if self._shuffle: + self._rng.shuffle(self._order) + self._position = 0 + sample_index = self._order[self._position] + self._position += 1 + return self._samples[sample_index] + + def state_dict(self) -> dict: + """Snapshot row order and position so resume continues the same stream.""" + return { + "rng_state": self._rng.getstate(), + "order": list(self._order), + "position": self._position, + } + + def load_state_dict(self, state_dict: dict) -> None: + """Restore state returned by `state_dict`.""" + self._rng.setstate(state_dict["rng_state"]) + self._order = list(state_dict["order"]) + self._position = state_dict["position"] + + +class DapoMathDataset(_CyclingDataset): + """Provides filtered DAPO-Math problems in the original `Answer:` format.""" + + @dataclass(kw_only=True, slots=True) + class Config(Configurable.Config): + repo_id: str = "hamishivi/DAPO-Math-17k-Processed_filtered" + split: str = "train" + seed: int = 42 + shuffle: bool = True + + def __init__(self, config: Config) -> None: + dataset = load_dataset(config.repo_id, split=config.split) + samples: list[DapoMathSample] = [] + for row in dataset: + prompt_messages = row["source_prompt"] + if len(prompt_messages) != 1 or prompt_messages[0]["role"] != "user": + raise ValueError("DAPO-Math rows must contain exactly one user prompt") + samples.append( + DapoMathSample( + prompt=prompt_messages[0]["content"], + ground_truth=str(row["ground_truth"]), + ) + ) + super().__init__(samples, seed=config.seed, shuffle=config.shuffle) + + +class AIME2025Dataset(_CyclingDataset): + """Provides AIME 2025 I+II problems using the DAPO answer format.""" + + @dataclass(kw_only=True, slots=True) + class Config(Configurable.Config): + repo_id: str = "opencompass/AIME2025" + subsets: tuple[str, ...] = ("AIME2025-I", "AIME2025-II") + split: str = "test" + seed: int = 99 + shuffle: bool = False + num_samples: int = 30 + + def __init__(self, config: Config) -> None: + dataset = concatenate_datasets( + [ + load_dataset(config.repo_id, subset, split=config.split) + for subset in config.subsets + ] + ).select(range(config.num_samples)) + samples = [ + DapoMathSample( + prompt=_AIME_PROMPT_TEMPLATE.format(problem=row["question"]), + ground_truth=str(row["answer"]), + ) + for row in dataset + ] + super().__init__(samples, seed=config.seed, shuffle=config.shuffle) diff --git a/torchtitan/experiments/rl/examples/dapo_math/env.py b/torchtitan/experiments/rl/examples/dapo_math/env.py new file mode 100644 index 00000000000..ec2ea858778 --- /dev/null +++ b/torchtitan/experiments/rl/examples/dapo_math/env.py @@ -0,0 +1,41 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +from dataclasses import dataclass + +from renderers import Message + +from torchtitan.experiments.rl.environment import ( + MessageEnv, + MessageEnvInitOutput, + MessageEnvStepOutput, +) +from torchtitan.experiments.rl.examples.dapo_math.data import DapoMathSample + + +class DapoMathEnv(MessageEnv): + """Single-turn message environment for a verifiable math problem.""" + + @dataclass(kw_only=True, slots=True) + class Config(MessageEnv.Config): + pass + + def __init__(self, config: Config, *, env_input: DapoMathSample) -> None: + del config + self._prompt = env_input.prompt + + async def init(self) -> MessageEnvInitOutput: + """Return the problem as one user message.""" + return MessageEnvInitOutput( + init_prompt_messages=[{"role": "user", "content": self._prompt}] + ) + + async def step(self, completion_message: Message) -> MessageEnvStepOutput: + """End after the model's first response; the rubric scores its final answer.""" + del completion_message + return MessageEnvStepOutput(done=True) diff --git a/torchtitan/experiments/rl/examples/dapo_math/requirements.txt b/torchtitan/experiments/rl/examples/dapo_math/requirements.txt new file mode 100644 index 00000000000..4ca03dae36f --- /dev/null +++ b/torchtitan/experiments/rl/examples/dapo_math/requirements.txt @@ -0,0 +1 @@ +math-verify==0.9.0 diff --git a/torchtitan/experiments/rl/examples/dapo_math/rollouter.py b/torchtitan/experiments/rl/examples/dapo_math/rollouter.py new file mode 100644 index 00000000000..bea6963fc7a --- /dev/null +++ b/torchtitan/experiments/rl/examples/dapo_math/rollouter.py @@ -0,0 +1,55 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +from dataclasses import dataclass, field + +from torchtitan.experiments.rl.environment import TokenEnv +from torchtitan.experiments.rl.examples.dapo_math.data import ( + AIME2025Dataset, + DapoMathDataset, +) +from torchtitan.experiments.rl.examples.dapo_math.env import DapoMathEnv +from torchtitan.experiments.rl.examples.dapo_math.rubric import RewardMathVerify +from torchtitan.experiments.rl.rollout.advantage import AdvantageEstimator +from torchtitan.experiments.rl.rollout.rollouter import Rollouter +from torchtitan.experiments.rl.rubrics import Rubric + + +class DapoMathRollouter(Rollouter): + """Single-turn math reasoning with a binary Math-Verify reward. + + Each sample produces one assistant solution; training uses DAPO-Math and + validation uses AIME 2025. + """ + + @dataclass(kw_only=True, slots=True) + class Config(Rollouter.Config): + train_dataset: DapoMathDataset.Config = field( + default_factory=DapoMathDataset.Config + ) + validation_dataset: AIME2025Dataset.Config = field( + default_factory=AIME2025Dataset.Config + ) + rubric: Rubric.Config = field( + default_factory=lambda: Rubric.Config( + reward_fns=[RewardMathVerify.Config(weight=1.0)], + error_reward=0.0, + ) + ) + message_env: DapoMathEnv.Config = field(default_factory=DapoMathEnv.Config) + token_env: TokenEnv.Config = field( + default_factory=lambda: TokenEnv.Config( + max_rollout_tokens=10240, + max_num_turns=1, + ) + ) + advantage: AdvantageEstimator.Config = field( + default_factory=lambda: AdvantageEstimator.Config( + should_std_normalize=False + ) + ) diff --git a/torchtitan/experiments/rl/examples/dapo_math/rubric.py b/torchtitan/experiments/rl/examples/dapo_math/rubric.py new file mode 100644 index 00000000000..27824b9eeb5 --- /dev/null +++ b/torchtitan/experiments/rl/examples/dapo_math/rubric.py @@ -0,0 +1,68 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +from dataclasses import dataclass + +from math_verify import LatexExtractionConfig, LatexNormalizationConfig, parse, verify +from math_verify.errors import TimeoutException + +from torchtitan.experiments.rl.examples.dapo_math.data import DapoMathSample +from torchtitan.experiments.rl.rollout import Rollout +from torchtitan.experiments.rl.rubrics import RewardFn + + +# Require an `Answer:` or `\boxed{}` marker so intermediate math is ignored. +_FINAL_ANSWER_EXTRACTION = [ + LatexExtractionConfig( + normalization_config=LatexNormalizationConfig(units=True), + boxed_match_priority=0, + try_extract_without_anchor=False, + ) +] + + +def score_math_response(response: str, ground_truth: str) -> float: + """Score an `Answer:` or `\\boxed{}` expression with Math-Verify. + + Args: + response: Model response containing a marked final answer. + ground_truth: Expected answer from the dataset. + + Example: + score_math_response("work\nAnswer: $34$", "34") # 1.0 + """ + try: + gold = parse(ground_truth) + prediction = parse( + response, + extraction_config=_FINAL_ANSWER_EXTRACTION, + extraction_mode="first_match", + ) + return float(bool(gold) and verify(gold, prediction)) + except (Exception, TimeoutException): + # Model output is untrusted; malformed LaTeX is an incorrect answer, not a + # training-loop failure. Math-Verify raises `TimeoutException` from BaseException. + return 0.0 + + +class RewardMathVerify(RewardFn): + """Binary reward for a mathematically equivalent final answer.""" + + @dataclass(kw_only=True, slots=True) + class Config(RewardFn.Config): + pass + + async def __call__(self, rollout: Rollout, env_input: DapoMathSample) -> float: + """Return 1 when Math-Verify equates the response and ground truth.""" + if not rollout.turns: + return 0.0 + completion_message = rollout.turns[-1].completion_message + response = ( + (completion_message.get("content") or "") if completion_message else "" + ) + return score_math_response(response, env_input.ground_truth) diff --git a/torchtitan/experiments/rl/tests/test_dapo_math.py b/torchtitan/experiments/rl/tests/test_dapo_math.py new file mode 100644 index 00000000000..1bedb64f9f2 --- /dev/null +++ b/torchtitan/experiments/rl/tests/test_dapo_math.py @@ -0,0 +1,120 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""CPU tests for the DAPO-Math dataset, environment, and rubric.""" + +from __future__ import annotations + +import asyncio + +from datasets import Dataset + +from torchtitan.experiments.rl.examples.dapo_math import ( + AIME2025Dataset, + DapoMathDataset, + DapoMathEnv, + DapoMathSample, + data as math_data, + RewardMathVerify, + score_math_response, +) +from torchtitan.experiments.rl.rollout import Rollout, RolloutStatus, RolloutTurn +from torchtitan.experiments.rl.types import RolloutTurnID + + +def _dapo_rows() -> list[dict]: + return [ + { + "source_prompt": [{"role": "user", "content": "problem 1"}], + "ground_truth": "34", + }, + { + "source_prompt": [{"role": "user", "content": "problem 2"}], + "ground_truth": "113", + }, + { + "source_prompt": [{"role": "user", "content": "problem 3"}], + "ground_truth": "7", + }, + ] + + +def test_dapo_dataset_is_deterministic_and_resumable(monkeypatch) -> None: + monkeypatch.setattr(math_data, "load_dataset", lambda *args, **kwargs: _dapo_rows()) + config = DapoMathDataset.Config(seed=7) + first = config.build() + second = config.build() + assert [next(first) for _ in range(3)] == [next(second) for _ in range(3)] + + checkpoint = first.state_dict() + expected = [next(first) for _ in range(3)] + resumed = config.build() + resumed.load_state_dict(checkpoint) + assert [next(resumed) for _ in range(3)] == expected + + +def test_aime_dataset_combines_both_subsets(monkeypatch) -> None: + def load_dataset(repo_id, subset, *, split): + del repo_id, split + answer = r"42^\circ" if subset == "AIME2025-I" else r"\boxed{42}" + return Dataset.from_list([{"question": f"{subset} question", "answer": answer}]) + + monkeypatch.setattr(math_data, "load_dataset", load_dataset) + dataset = AIME2025Dataset.Config(num_samples=2).build() + samples = [next(dataset), next(dataset)] + assert [sample.ground_truth for sample in samples] == [r"42^\circ", r"\boxed{42}"] + assert "AIME2025-I question" in samples[0].prompt + assert "AIME2025-II question" in samples[1].prompt + + +def test_aime_dataset_restarts_after_configured_num_samples(monkeypatch) -> None: + def load_dataset(repo_id, subset, *, split): + del repo_id, split + return Dataset.from_list([{"question": f"{subset} question", "answer": "42"}]) + + monkeypatch.setattr(math_data, "load_dataset", load_dataset) + dataset = AIME2025Dataset.Config(num_samples=1).build() + first = next(dataset) + assert next(dataset) == first + + +def test_env_is_single_turn() -> None: + env = DapoMathEnv.Config().build( + env_input=DapoMathSample(prompt="solve me", ground_truth="3"), + ) + initial = asyncio.run(env.init()) + assert initial.init_prompt_messages == [{"role": "user", "content": "solve me"}] + assert asyncio.run(env.step({"role": "assistant", "content": "Answer: 3"})).done + + +def _rollout(response: str) -> Rollout: + return Rollout( + group_id=0, + rollout_id=0, + status=RolloutStatus.COMPLETED, + turns=[ + RolloutTurn( + rollout_id=RolloutTurnID(group_id=0, rollout_id=0, turn_id=0), + prompt_token_ids=[1], + completion_token_ids=[2], + completion_logprobs=[-0.1], + completion_message={"role": "assistant", "content": response}, + ) + ], + ) + + +def test_math_verifier_requires_an_answer_marker() -> None: + assert score_math_response("work\nAnswer: $34$", "34") == 1.0 + assert score_math_response(r"work\n\boxed{34}", "34") == 1.0 + assert score_math_response("work mentions 34", "34") == 0.0 + + +def test_reward_handles_equivalent_latex_and_units() -> None: + reward = RewardMathVerify.Config().build() + sample = DapoMathSample(prompt="problem", ground_truth=r"336^\circ") + assert asyncio.run(reward(_rollout("work\nAnswer: $336$"), sample)) == 1.0 + assert asyncio.run(reward(_rollout("work\nAnswer: $335$"), sample)) == 0.0