GSM8K
Category: RL; training tokens 1.4M; inference tokens 1.2M
Introduction
GSM8K (Grade School Math 8K) is a dataset of 8,500 high-quality, linguistically diverse grade-school math word problems. It is designed for tasks that require solving basic math problems through multi-step reasoning. It has these properties:
- The problems usually require 2 to 8 reasoning steps.
- The solutions mainly involve basic arithmetic operations (+ - x /).
- A capable middle-school student should be able to solve each problem.
- The solutions are provided in natural language rather than pure mathematical expressions.
In LLM reinforcement learning, GSM8K is a standard benchmark for evaluating reasoning ability and one of the most common RL training scenarios.
Example item:
{
"question": "Natalia sold clips to 48 of her friends in April, and then she sold half as many clips in May. How many clips did Natalia sell altogether in April and May?",
"answer": "Natalia sold 48/2 = <<48/2=24>>24 clips in May.\nNatalia sold 48+24 = <<48+24=72>>72 clips altogether in April and May.\n#### 72"
}The RL objective for GSM8K is to improve the model's ability to solve multi-step math word problems by repeatedly attempting answers, receiving rule-based reward feedback, and updating its reasoning policy.
Environment
Install dependencies on any CPU-only machine with internet access:
pip install pytrio transformers modelscope datasets tqdm swanlab numpyswanlab is used to observe training curves. Log in locally before use. See SwanLab Quick Start.
Dataset
Download GSM8K into the training project's gsm8k/ directory:
modelscope download --dataset AI-ModelScope/gsm8k --local_dir ./gsm8kCode
Training and evaluation use about 1.4M training tokens and 1.2M inference tokens. This example uses Qwen/Qwen3.5-4B as the default base model. One reference run takes about 13 minutes, but actual runtime varies with queueing, sampling length, and sample count.
The API calls in this example are aligned with the API documentation:
- Use
ServiceClient.create_lora_training_client_async()to create the LoRA training client. - Use
TrainingClient.save_weights_and_get_sampling_client_async()inside the training loop to temporarily save the current LoRA weights and create a sampling client. This method does not take a name. - Use
SamplingClient.sample_async()to submit sampling requests and receiveSampleResponseobjects. - Use
target_tokens,logprobs, andadvantagesinDatum.loss_fn_inputsto buildimportance_samplingtraining samples. - Use
TrainingClient.forward_backward_async(..., "importance_sampling")to accumulate gradients, then calloptim_step_async()to update parameters. - Save the final LoRA weights with
save_weights_for_sampler_async(name=...), then load them for evaluation withcreate_sampling_client_async(base_model=..., model_path=...).
Run the following code to start training:
"""TRIO + GSM8K importance-sampling RL fine-tuning tutorial example.
Core flow:
1. At each step, temporarily create a sampler from the current LoRA weights.
2. The sampler asynchronously samples completions for the current batch and returns completion, sampling logprob, and reward.
3. Rewards for the same question are normalized into group-relative advantages and converted into trio.Datum objects.
4. forward_backward(..., "importance_sampling") computes gradients, then optim_step updates weights.
"""
import argparse
import asyncio
import math
import re
import time
import numpy as np
import pytrio as trio
import swanlab
from datasets import load_dataset
from tqdm.asyncio import tqdm_asyncio
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="TRIO on-policy RL fine-tuning example for GSM8K.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument("--base-model", default="Qwen/Qwen3.5-4B", help="TRIO trainable base model")
parser.add_argument("--dataset-path", default="./gsm8k", help="Local GSM8K dataset path")
parser.add_argument("--dataset-config", default="main", help="Dataset config name for datasets.load_dataset")
parser.add_argument("--lora-rank", type=int, default=32, help="LoRA rank")
parser.add_argument("--epochs", type=int, default=1, help="Number of passes over the training subset")
parser.add_argument("--train-samples", type=int, default=512, help="Number of GSM8K samples for training")
parser.add_argument("--eval-samples", type=int, default=256, help="Number of GSM8K samples for final evaluation")
parser.add_argument("--prompt-batch-size", type=int, default=8, help="Number of questions sampled per RL step")
parser.add_argument("--num-samples-per-prompt", type=int, default=4, help="Number of completions sampled per question")
parser.add_argument("--max-tokens", type=int, default=512, help="Maximum generated tokens per completion")
parser.add_argument("--temperature", type=float, default=0.7, help="Sampling temperature for training")
parser.add_argument("--learning-rate", type=float, default=1e-5, help="AdamW learning rate")
parser.add_argument("--seed", type=int, default=None, help="Random seed. None means no fixed seed")
parser.add_argument("--eval", dest="eval_model_path", default=None, help="Evaluation-only mode with a sampler path")
parser.add_argument("--checkpoint-prefix", default="rl-gsm8k", help="Prefix used when saving TRIO sampler weights")
parser.add_argument("--swanlab-project", default="GSM8K-WITH-TRIO", help="SwanLab project name")
parser.add_argument("--swanlab-experiment", default="rl-gsm8k", help="SwanLab experiment name")
return parser.parse_args()
def make_prompt(question: str) -> str:
return (
f"Question: {question}\n"
"Let's think step by step. Put your final numeric answer after '#### '.\n"
"Answer:"
)
def gold_answer(answer: str) -> float:
return float(answer.split("####")[-1].strip().replace(",", ""))
def parse_model_answer(text: str) -> float | None:
"""Prefer the answer after ####; fall back to the last number if missing."""
ANSWER_RE = re.compile(r"####\s*(-?\d+(?:\.\d+)?)")
NUMBER_RE = re.compile(r"-?\d+(?:\.\d+)?")
clean = text.replace(",", "")
match = ANSWER_RE.search(clean)
if match:
return float(match.group(1))
numbers = NUMBER_RE.findall(clean)
return float(numbers[-1]) if numbers else None
def reward_fn(text: str, gold: float) -> float:
# Rule-based reward for teaching: correct answers get positive reward; wrong or unparsable answers are penalized.
pred = parse_model_answer(text)
if pred is None:
return -1.0
return 1.0 if abs(pred - gold) < 1e-6 else -0.5
def group_advantages(rewards: list[float]) -> list[float]:
"""Normalize rewards within the same question to get group-relative advantages."""
if not rewards:
return []
mean = float(np.mean(rewards))
std = float(np.std(rewards))
return [(reward - mean) / (std + 1e-8) for reward in rewards]
def normalize_logprobs(logprobs: list[float | None]) -> list[float]:
"""Normalize sampling logprobs into floats so None values do not affect TensorData construction."""
return [0.0 if value is None else float(value) for value in logprobs]
def to_numpy(value) -> np.ndarray:
"""Support PyTrio TensorData, numpy arrays, and plain lists."""
if hasattr(value, "to_numpy"):
return np.asarray(value.to_numpy()).reshape(-1)
if hasattr(value, "tolist"):
return np.asarray(value.tolist()).reshape(-1)
return np.asarray(value).reshape(-1)
def metric_loss(metrics: dict[str, float], token_count: int) -> float | None:
"""Read loss from server metrics. Return None if the metric is unavailable."""
for key in ("loss:sum", "loss_sum", "loss/total", "loss_total"):
if key in metrics:
return float(metrics[key]) / max(token_count, 1)
for key in ("loss", "loss:mean", "loss_mean", "loss/mean"):
if key in metrics:
return float(metrics[key])
return None
def compute_importance_sampling_loss(
datums: list[trio.Datum],
loss_fn_outputs: list[dict],
) -> float:
"""Compute the importance-sampling objective locally from returned current logprobs for logging."""
losses = []
for datum, output in zip(datums, loss_fn_outputs, strict=True):
if "logprobs" not in output:
raise KeyError(f"forward_backward output has no logprobs; output keys: {list(output)}")
current_logprobs = to_numpy(output["logprobs"]).astype(np.float32)
old_logprobs = to_numpy(datum.loss_fn_inputs["logprobs"]).astype(np.float32)
advantages = to_numpy(datum.loss_fn_inputs["advantages"]).astype(np.float32)
length = min(len(current_logprobs), len(old_logprobs), len(advantages))
if length == 0:
continue
ratios = np.exp(current_logprobs[:length] - old_logprobs[:length])
token_mask = advantages[:length] != 0
if np.any(token_mask):
losses.append(-(ratios[token_mask] * advantages[:length][token_mask]))
if not losses:
return 0.0
return float(np.concatenate(losses).mean())
def make_datum(
prompt_tokens: list[int],
completion_tokens: list[int],
completion_logprobs: list[float | None],
advantage: float,
) -> trio.Datum | None:
"""Convert one completion into the Datum format required by TRIO importance_sampling loss."""
if not completion_tokens:
return None
tokens = prompt_tokens + completion_tokens
# The prompt is context only and does not contribute to loss; completion tokens use the advantage.
weights = ([0.0] * len(prompt_tokens) + [1.0] * len(completion_tokens))
advantages = [advantage * weight for weight in weights]
# importance_sampling needs logprobs from the old sampling policy. Prompt positions are padded with 0 and masked by advantages.
completion_logprobs = normalize_logprobs(completion_logprobs)
old_logprobs = ([0.0] * len(prompt_tokens) + completion_logprobs)[: len(tokens)]
old_logprobs += [0.0] * (len(tokens) - len(old_logprobs))
# model_input, target_tokens, logprobs, and advantages are shifted by one and must have the same length.
return trio.Datum(
model_input=trio.ModelInput.from_ints(tokens=tokens[:-1]),
loss_fn_inputs={
"target_tokens": tokens[1:],
"logprobs": old_logprobs[1:],
"advantages": advantages[1:],
},
)
def iter_batches(dataset: list[dict], batch_size: int, epochs: int):
for epoch in range(epochs):
for start in range(0, len(dataset), batch_size):
yield epoch, start, dataset[start : start + batch_size]
async def sample_one_question(sampler, tokenizer, item: dict, args: argparse.Namespace) -> dict:
"""Sample multiple completions for one question and compute advantages within the question group."""
prompt_tokens = tokenizer.encode(make_prompt(item["question"]), add_special_tokens=True)
sample_result = await sampler.sample_async(
prompt=trio.ModelInput.from_ints(prompt_tokens),
sampling_params=trio.SamplingParams(
max_tokens=args.max_tokens,
temperature=args.temperature,
seed=args.seed,
),
num_samples=args.num_samples_per_prompt,
)
gold = gold_answer(item["answer"])
completions = []
completion_lens = []
rewards = []
for sequence in sample_result.sequences:
completion_tokens = list(sequence.tokens)
reward = reward_fn(sequence.text, gold)
pred = parse_model_answer(sequence.text)
is_correct = pred is not None and abs(pred - gold) < 1e-6
completions.append((completion_tokens, sequence.logprobs, is_correct))
completion_lens.append(len(completion_tokens))
rewards.append(reward)
advantages = group_advantages(rewards)
corrects = []
datums = []
for (completion_tokens, logprobs, is_correct), advantage in zip(completions, advantages):
datum = make_datum(prompt_tokens, completion_tokens, logprobs, advantage)
if datum is not None:
datums.append(datum)
corrects.append(is_correct)
correct = sum(corrects)
return {
"datums": datums,
"rewards": rewards,
"advantages": advantages,
"correct": correct,
"comp_len": completion_lens,
}
async def collect_rollouts(sampler, tokenizer, batch: list[dict], args: argparse.Namespace):
"""Sample one prompt batch concurrently and return training Datum objects plus logging metrics."""
# Questions in a batch are independent, so sampling requests can be submitted concurrently.
results = await asyncio.gather(
*(sample_one_question(sampler, tokenizer, item, args) for item in batch)
)
datums = [datum for result in results for datum in result["datums"]]
rewards = [reward for result in results for reward in result["rewards"]]
advantages = [adv for result in results for adv in result["advantages"]]
correct = sum(result["correct"] for result in results)
completion_lens = [comp_len for result in results for comp_len in result["comp_len"]]
if not datums:
print("No valid datums, skip this batch")
return [], {}
return datums, {
"reward_mean": float(np.mean(rewards)),
"reward_std": float(np.std(rewards)),
"advantage_std": float(np.std(advantages)),
"accuracy": correct / len(datums),
"completion_len_avg": float(np.mean(completion_lens)),
"completion_len_std": float(np.std(completion_lens)),
"batch_train_tokens": sum(completion_lens),
}
async def train(
training_client,
tokenizer,
train_dataset: list[dict],
args: argparse.Namespace,
) -> int:
total_steps = args.epochs * math.ceil(len(train_dataset) / args.prompt_batch_size)
# Cosine learning-rate schedule.
lr_schedule = lambda step: args.learning_rate * 0.5 * (1 + math.cos(math.pi * step / total_steps))
print("Start on-policy importance sampling RL training")
for step, (epoch, batch_start, batch) in enumerate(
iter_batches(train_dataset, args.prompt_batch_size, args.epochs)
):
loop_start_time = time.time()
# save_weights_and_get_sampling_client_async saves current LoRA weights to a temporary archive
# and returns a SamplingClient loaded with those weights. Per the API docs, it takes no name.
sampler = await training_client.save_weights_and_get_sampling_client_async()
datums, rollout_stats = await collect_rollouts(sampler, tokenizer, batch, args)
if not datums:
continue
fwdbwd_future = await training_client.forward_backward_async(datums, "importance_sampling")
learning_rate = lr_schedule(step)
optim_future = await training_client.optim_step_async(
trio.AdamParams(learning_rate=learning_rate)
)
fwdbwd_result, _ = await asyncio.gather(fwdbwd_future, optim_future)
loss = metric_loss(fwdbwd_result.metrics, rollout_stats["batch_train_tokens"])
if loss is None:
loss = compute_importance_sampling_loss(datums, fwdbwd_result.loss_fn_outputs)
loop_used_time = time.time() - loop_start_time
metrics = {
"train/loss": loss,
"train/learning_rate": learning_rate,
**{f"rollout/{key}": value for key, value in rollout_stats.items()},
"epoch": epoch,
"batch_start": batch_start,
"loop_time": loop_used_time,
}
metrics.update({f"trainer/{key}": value for key, value in fwdbwd_result.metrics.items()})
swanlab.log(metrics, step=step)
print(
f"Step {step + 1}/{total_steps} | Epoch {epoch + 1} | "
f"Reward {rollout_stats['reward_mean']:.3f} | "
f"Acc {rollout_stats['accuracy']:.3f} | "
f"Batch {len(datums)} | "
f"Learning Rate {learning_rate:.3e} | "
f"Loss {loss:.4f} | "
f"Loop Time {loop_used_time:.2f}"
)
return total_steps
async def evaluate(
name: str,
sampler,
tokenizer,
eval_dataset: list[dict],
args: argparse.Namespace,
) -> dict:
"""Evaluate one model by sampling concurrently, parsing answers, and computing accuracy."""
print(f"Evaluating {name} model...")
params = trio.SamplingParams(max_tokens=args.max_tokens, temperature=0.0, seed=args.seed)
examples = []
prompts = []
for item in eval_dataset:
gold = gold_answer(item["answer"])
prompt = trio.ModelInput.from_ints(
tokenizer.encode(make_prompt(item["question"]), add_special_tokens=True)
)
examples.append((item["question"], gold))
prompts.append(prompt)
# sample_async directly returns SampleResponse; tqdm_asyncio.gather handles concurrent sampling and progress display.
sample_results = await tqdm_asyncio.gather(
*(sampler.sample_async(prompt=prompt, sampling_params=params, num_samples=1) for prompt in prompts),
desc="Evaluating",
)
correct = 0
print_top_k = 3
for (question, gold), sample_result in zip(examples, sample_results):
text = sample_result.sequences[0].text
pred = parse_model_answer(text)
is_correct = pred is not None and abs(pred - gold) < 1e-6
correct += is_correct
if print_top_k > 0:
print("=" * 80)
print(f"Model: {name}")
print(f"Q: {question}")
print(f"Gold: {gold}")
print(f"Pred: {repr(text.strip())} -> {pred}")
print(f"Correct: {is_correct}")
print_top_k -= 1
total = len(eval_dataset)
metrics = {
"accuracy": correct / max(total, 1),
"correct": correct,
"total": total,
}
print("=" * 80)
print(f"{name} Accuracy: {metrics['accuracy']:.4f} ({correct}/{total})")
return metrics
async def main():
# Parse command-line arguments.
args = parse_args()
# Connect to the TRIO service.
print("Connecting to TRIO service...")
service_client = trio.ServiceClient()
# Load the GSM8K dataset.
print("Loading GSM8K dataset...")
gsm8k = load_dataset(args.dataset_path, args.dataset_config)
eval_dataset = list(gsm8k["test"])[: args.eval_samples]
# Evaluation-only mode: skip training and evaluate the specified model directly.
if args.eval_model_path:
eval_sampler = await service_client.create_sampling_client_async(
base_model=args.base_model,
model_path=args.eval_model_path,
)
await evaluate("eval", eval_sampler, eval_sampler.get_tokenizer(), eval_dataset, args)
return
# Create a LoRA training client.
training_client = await service_client.create_lora_training_client_async(
base_model=args.base_model,
rank=args.lora_rank,
seed=args.seed,
)
tokenizer = training_client.get_tokenizer()
train_dataset = list(gsm8k["train"])[: args.train_samples]
# Initialize SwanLab experiment tracking.
swanlab.init(
project=args.swanlab_project,
experiment_name=args.swanlab_experiment,
config=vars(args) | {"loss_fn": "importance_sampling"},
)
# Run RL training.
total_steps = await train(training_client, tokenizer, train_dataset, args)
# Save the final model.
print("Saving final model...")
rl_sampler_future = await training_client.save_weights_for_sampler_async(name=f"{args.checkpoint_prefix}-final")
rl_sampler_result = await rl_sampler_future
print(f"Final model saved to: {rl_sampler_result.path}")
# After training, evaluate both the base model and the RL-tuned model.
print("Start Evaluation on GSM8K Test Set")
base_sampler = await service_client.create_sampling_client_async(
base_model=args.base_model
)
rl_sampler = await training_client.create_sampling_client_async(
model_path=rl_sampler_result.path,
)
base_metrics = await evaluate("Base Model", base_sampler, tokenizer, eval_dataset, args)
rl_metrics = await evaluate("RL Model", rl_sampler, tokenizer, eval_dataset, args)
# Log final evaluation metrics to SwanLab.
swanlab.log({
"eval/base_accuracy": base_metrics["accuracy"],
"eval/rl_accuracy": rl_metrics["accuracy"],
"eval/base_correct": base_metrics["correct"],
"eval/rl_correct": rl_metrics["correct"],
"eval/total": base_metrics["total"],
}, step=total_steps)
if __name__ == "__main__":
asyncio.run(main())Training Results
In one reference run, after 1 epoch of training, the base model reached 87.1% accuracy on the test set, while the RL-tuned model reached 94.1%, substantially improving math problem-solving performance.
...
Final model saved to: trio004:48ltn0x3j9/2pbnakmve6fe/weights/rl-gsm8k-final
Start Evaluation on GSM8K Test Set
================================================================================
Base Model Accuracy: 0.8711 (223/256)
RL Model Accuracy: 0.9414 (241/256)
Chat-Huanhuan
Category: SFT; training tokens 0.6M
GRPO
Category: RL; dataset: GSM8K; training tokens 0.132M; prefill tokens 0.1M; sample tokens 0.311M; implementations: sync / async