Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions torchtitan/experiments/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
# RL examples own a per-example config_registry under rl/examples/<name>;
# listed here so `--module <name>` resolves (see ConfigManager).
"alphabet_sort",
"dapo_math",
"search_r1",
]
)
102 changes: 102 additions & 0 deletions torchtitan/experiments/rl/examples/dapo_math/README.md
Original file line number Diff line number Diff line change
@@ -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.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

can we also put a validation set accuracy figure in the readme so people can be more convinced

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

yes, i ahve to rerun it. I left as a todo at the bottom of the readme


## 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)
Comment on lines +100 to +102

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Curves in these pictures would be easy to get outdated and not reproducible on later commit. Maybe only add them in PR summary, or add a commit pin to the results.

27 changes: 27 additions & 0 deletions torchtitan/experiments/rl/examples/dapo_math/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
]
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
161 changes: 161 additions & 0 deletions torchtitan/experiments/rl/examples/dapo_math/config_registry.py
Original file line number Diff line number Diff line change
@@ -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",
)
Loading
Loading