RL Training with Tinker (using Unified Trainer)
This example shows how to train a solver‑judge RL workflow with the Tinker backend in rLLM, using Tinker's hosted GPU service and the Unified Trainer.
Overview
With this example you will:
- Train a solver‑judge workflow for the Countdown task using the Tinker backend
- Use the Unified Trainer (
rllm.experimental.unified_trainer.AgentTrainer) with the unified Hydra config system
Under the hood, rLLM integrates with Tinker as:
- Rollout backend: sampling and logprob computation happen on Tinker's GPU service
- Policy trainer: LoRA adapters are optimized remotely via Tinker training clients
- Checkpoint manager: checkpoints are stored and resumed via Tinker model IDs
Setup
Install dependencies
Configure authentication
Set your API keys:
export TINKER_API_KEY=your_api_key_here
export WANDB_API_KEY=your_wandb_key_here # optional, for W&B logging
You can obtain a Tinker API key from the Tinker console.
Unified configuration system
This example uses the unified Hydra config introduced with the experimental unified trainer. Configuration is organized into:
- Backend-agnostic rLLM configs (
rllm/experimental/config/rllm/base.yaml): core training settings shared across all backends - Tinker-specific configs (
rllm/experimental/config/rllm/backend/tinker.yaml): Tinker service, model, sampling, and data settings - Unified entry point (
rllm/experimental/config/unified.yaml): combines all configs
Key options you may want to tune:
| Config path | Description |
|---|---|
model.name |
Base model to fine‑tune (e.g. Qwen/Qwen3-8B) |
model.lora_rank |
LoRA rank |
training.group_size |
Number of trajectories per prompt (GRPO group size) |
data.max_prompt_length / data.max_response_length |
Context and generation lengths |
rllm.trainer.total_epochs / rllm.trainer.total_batches |
Training budget |
rllm.trainer.logger |
Logging backends (console, wandb, tensorboard) |
rllm.algorithm.adv_estimator |
Advantage estimator (grpo, reinforce, rloo, etc.) |
Backend-specific keys (under model.*, training.*, sampling.*, data.*) are forwarded into the common rllm.* namespace automatically. See the rLLM and Backend Config docs for details.
Solver‑Judge RL Training with Tinker
This example trains a multi‑agent solver‑judge workflow on the Countdown task using the Tinker backend and the Unified Trainer.
1. Prepare Countdown dataset
First download and register the Countdown dataset:
This will:
- Load
Jiayi-Pan/Countdown-Tasks-3to4from HuggingFace - Convert each example into a math‑style word problem
- Register multiple splits (train, test, stage2, stage3) under the
countdownkey
Dataset preparation:
import random
from datasets import load_dataset
from rllm.data.dataset import DatasetRegistry
def prepare_countdown_data():
"""
Prepare the countdown task dataset from HuggingFace.
Take 1024 examples as test set, remaining as training set.
Also create stage 2 and stage 3 training sets with 50k examples each.
"""
# Load the countdown dataset
dataset = load_dataset("Jiayi-Pan/Countdown-Tasks-3to4", split="train")
# Split dataset: 1024 examples for test, rest for training
test_size = 1024
total_size = len(dataset)
# Create train/test split
test_dataset = dataset.select(range(test_size))
train_dataset = dataset.select(range(test_size, total_size))
def preprocess_fn(example, idx):
"""
Convert countdown task format to math problem format.
Example: target=98, nums=[44, 19, 35] becomes a math word problem.
"""
target = example["target"]
nums = example["nums"]
# Format as a math problem
nums_str = ", ".join(map(str, nums))
question = f"Using the numbers {nums_str}, find a way to reach the target number {target}. You can use basic arithmetic operations (+, -, *, /) and each number can only be used once. Show your step-by-step calculation and output the final answer within <answer>...</answer>, for example <answer> (1 + 2) / 3 </answer>."
return {
"question": question,
"ground_truth": str(target),
"data_source": "countdown",
"target": target,
"nums": nums,
}
# Apply preprocessing
train_dataset = train_dataset.map(preprocess_fn, with_indices=True)
test_dataset = test_dataset.map(preprocess_fn, with_indices=True)
# Create stage 2 and stage 3 training datasets
train_size = len(train_dataset)
stage_size = 50000
# Ensure we have enough data for both stages
if train_size < 2 * stage_size:
print(f"Warning: Training set has only {train_size} examples, but need {2 * stage_size} for both stages")
stage_size = min(stage_size, train_size // 2)
# Shuffle and select indices for stage 2 and stage 3
all_indices = list(range(train_size))
random.shuffle(all_indices)
stage2_indices = all_indices[:stage_size]
stage3_indices = all_indices[stage_size : 2 * stage_size]
# Create stage datasets
stage2_dataset = train_dataset.select(stage2_indices)
stage3_dataset = train_dataset.select(stage3_indices)
# Register datasets
train_dataset = DatasetRegistry.register_dataset("countdown", train_dataset, "train")
test_dataset = DatasetRegistry.register_dataset("countdown", test_dataset, "test")
stage2_dataset = DatasetRegistry.register_dataset("countdown", stage2_dataset, "stage2_train")
stage3_dataset = DatasetRegistry.register_dataset("countdown", stage3_dataset, "stage3_train")
print(f"Train dataset size: {len(train_dataset)}")
print(f"Test dataset size: {len(test_dataset)}")
print(f"Stage 2 train dataset size: {len(stage2_dataset)}")
print(f"Stage 3 train dataset size: {len(stage3_dataset)}")
return train_dataset, test_dataset, stage2_dataset, stage3_dataset
if __name__ == "__main__":
train_dataset, test_dataset, stage2_dataset, stage3_dataset = prepare_countdown_data()
print("Train dataset path:", train_dataset.get_data_path())
print("Test dataset path:", test_dataset.get_data_path())
print("Stage 2 train dataset path:", stage2_dataset.get_data_path())
print("Stage 3 train dataset path:", stage3_dataset.get_data_path())
# Print a sample
print("\nSample train example:")
print(train_dataset[0])
print("\nSample stage 2 train example:")
print(stage2_dataset[0])
print("\nSample stage 3 train example:")
print(stage3_dataset[0])
2. Training entrypoint
The training entrypoint uses AgentTrainer from the experimental unified trainer:
"""
Adapted from `examples.solver_judge_tinker.train_solver_judge_flow_tinker` to
test the unified trainer with Tinker backend.
"""
import hydra
from examples.solver_judge.solver_judge_flow import SolverJudgeWorkflow
from rllm.data.dataset import DatasetRegistry
from rllm.experimental.common.config import rLLMAdvantageEstimator
from rllm.experimental.unified_trainer import AgentTrainer
from rllm.rewards.countdown_reward import countdown_reward_fn
@hydra.main(config_path="pkg://rllm.experimental.config", config_name="unified", version_base=None)
def main(config):
train_dataset = DatasetRegistry.load_dataset("countdown", "train")
test_dataset = DatasetRegistry.load_dataset("countdown", "test")
traj_group_adv_estimator_map = {
"solver": rLLMAdvantageEstimator.GRPO,
"judge": rLLMAdvantageEstimator.REINFORCE,
}
trainer = AgentTrainer(
workflow_class=SolverJudgeWorkflow,
workflow_args={
"n_solutions": 2,
"reward_function": countdown_reward_fn,
},
config=config,
train_dataset=train_dataset,
val_dataset=test_dataset,
backend="tinker",
traj_group_adv_estimator_map=traj_group_adv_estimator_map,
)
trainer.train()
if __name__ == "__main__":
main()
Key differences from the legacy trainer:
- Uses
rllm.experimental.unified_trainer.AgentTrainerinstead ofrllm.trainer.AgentTrainer - Config is loaded from
rllm.experimental.configwithconfig_name="unified" - Supports per‑role advantage estimator mapping via
traj_group_adv_estimator_map(e.g. GRPO for solver, REINFORCE for judge)
3. Train solver‑judge workflow with Tinker
Run training with the following script:
export TINKER_API_KEY=YOUR-TINKER-KEY
export WANDB_API_KEY=YOUR-WANDB-KEY # optional
MODEL_PATH=Qwen/Qwen3-4B-Instruct-2507
MODEL_LORA_RANK=32
# structure runname as <model_path>_<lora_rank>_<date>_<time>
date_str=$(date +%Y-%m-%d)
time_str=$(date +%H-%M-%S)
run_name=TINKER_${MODEL_PATH}_${MODEL_LORA_RANK}_${date_str}_${time_str}
local_dir=/path/to/your/local/dir
python3 -m rllm.experimental.test_examples.test_tinker_solver_judge \
rllm/backend=tinker \
rllm.compact_filtering.enable=False \
rllm.algorithm.adv_estimator=grpo \
rllm.algorithm.use_rllm=true \
rllm.algorithm.norm_adv_by_std_in_grpo=true \
rllm.trainer.total_batches=100 \
rllm.trainer.total_epochs=1 \
rllm.trainer.logger=['wandb'] \
rllm.trainer.project_name='tinker-solver-judge' \
rllm.trainer.experiment_name=$run_name \
rllm.trainer.val_before_train=true \
rllm.trainer.test_freq=20 \
rllm.trainer.save_freq=20 \
model.name=$MODEL_PATH \
model.lora_rank=$MODEL_LORA_RANK \
training.group_size=5 \
training.learning_rate=4e-5 \
training.default_local_dir=$local_dir \
sampling.train.temperature=1.0 \
sampling.train.top_p=1.0 \
data.max_prompt_length=2048 \
data.max_response_length=1024 \
data.train_batch_size=32 \
data.val_batch_size=512
This will:
- Fine‑tune
Qwen/Qwen3-4B-Instruct-2507with LoRA (rank 32) - Use rLLM‑native advantage computation with per‑role estimators (GRPO for solver, REINFORCE for judge)
- Normalize advantages by standard deviation for stability
- Run validation before training and every 20 batches
- Log training metrics to Weights & Biases
You can customize training via Hydra CLI overrides. Note that:
- rLLM-level settings use the
rllm.*prefix (e.g.rllm.algorithm.adv_estimator,rllm.trainer.total_batches) - Tinker backend settings use their native keys (e.g.
model.name,training.group_size,sampling.train.temperature)
For example, to switch to a larger model with a smaller LoRA rank:
python3 -m rllm.experimental.test_examples.test_tinker_solver_judge \
rllm/backend=tinker \
model.name=Qwen/Qwen3-8B \
model.lora_rank=16 \
training.group_size=8 \
data.train_batch_size=32 \
rllm.trainer.total_batches=200 \
rllm.trainer.logger=['console','wandb'] \
rllm.trainer.project_name='solver-judge-tinker' \
rllm.trainer.experiment_name='countdown-grpo-qwen3-8b'
4. Run the workflow with Tinker rollout engine (optional)
For interactive evaluation (no training step), you can run the Countdown solver‑judge workflow directly using Tinker for sampling:
import asyncio
import json
import os
# Import countdown-specific modules
import sys
from copy import deepcopy
import tinker
from solver_judge_flow import SolverJudgeWorkflow
from transformers import AutoTokenizer
from rllm.data.dataset import DatasetRegistry
from rllm.engine.agent_workflow_engine import AgentWorkflowEngine
from rllm.engine.rollout.tinker_engine import TinkerEngine
from rllm.rewards.countdown_reward import countdown_reward_fn
sys.path.append(os.path.join(os.path.dirname(__file__), "..", "countdown"))
def load_data(n=1):
"""Load countdown data using the Dataset interface."""
dataset = DatasetRegistry.load_dataset("countdown", "test")
if dataset is None:
print("Dataset not found, preparing dataset...")
from prepare_countdown_data import prepare_countdown_data
_, dataset, _, _ = prepare_countdown_data()
data = []
for idx, example in enumerate(dataset):
processed = process_countdown_fn(example, idx)
for i in range(n):
data.append(deepcopy(processed))
return data
def process_countdown_fn(example, idx):
"""Process countdown example into the expected format."""
question = example["question"]
target = example["target"]
nums = example["nums"]
# Create ground truth in the format expected by countdown_reward_fn
ground_truth = {"target": target, "numbers": nums}
task = {"question": question, "ground_truth": ground_truth, "idx": idx, "data_source": "countdown", "target": target, "nums": nums}
return task
def evaluate_results(results):
"""Evaluate the results and compute pass@k metrics."""
from collections import defaultdict
# Create a map to store correct answers per problem
problem_correct_map = defaultdict(int)
problem_total_map = defaultdict(int)
# Count correct answers for each problem
for episode in results:
problem = episode.task["question"]
# Use the episode-level is_correct flag set by the workflow
is_correct = episode.is_correct
problem_correct_map[problem] += int(is_correct)
problem_total_map[problem] += 1
# Calculate pass@1 and pass@k
k = max(problem_total_map.values()) if problem_total_map else 1
total_problems = len(problem_correct_map)
if total_problems > 0:
pass_at_1 = sum(problem_correct_map.values()) / sum(problem_total_map.values())
pass_at_k = sum(1 for problem, correct in problem_correct_map.items() if correct > 0) / total_problems
else:
pass_at_1 = 0.0
pass_at_k = 0.0
print("Total unique problems:", total_problems)
print("Average Pass@1 Accuracy:", pass_at_1)
print(f"Average Pass@{k} Accuracy:", pass_at_k)
if __name__ == "__main__":
import os
os.environ["TOKENIZERS_PARALLELISM"] = "true"
# Configuration
n_parallel_tasks = 4
n_solutions = 2 # Number of solutions to generate per problem
model_name = "Qwen/Qwen3-8B"
service_client = tinker.ServiceClient(base_url=None)
tokenizer = AutoTokenizer.from_pretrained(model_name)
rollout_engine = TinkerEngine(
model_name=model_name,
tokenizer=tokenizer,
service_client=service_client,
max_prompt_length=2048,
max_response_length=1024,
sampling_params={"temperature": 0.6, "top_p": 0.95},
)
training_client = service_client.create_lora_training_client(
base_model=model_name,
rank=4,
)
sampler_future = training_client.save_weights_for_sampler(name="000000")
sampler_result = sampler_future.result()
sampling_client = training_client.create_sampling_client(sampler_result.path)
rollout_engine.set_sampling_client(sampling_client)
engine = AgentWorkflowEngine(
workflow_cls=SolverJudgeWorkflow,
workflow_args={
"n_solutions": n_solutions,
"reward_function": countdown_reward_fn,
},
rollout_engine=rollout_engine,
config=None,
n_parallel_tasks=n_parallel_tasks,
retry_limit=1,
)
# Load countdown tasks
tasks = load_data(n=1)
print(f"Loaded {len(tasks)} countdown tasks")
tasks = tasks[:4]
results = asyncio.run(engine.execute_tasks(tasks))
import pdb
pdb.set_trace()
print(results[1])
# Evaluate results (rewards are already assigned in the workflow)
print("Evaluating results...")
evaluate_results(results)
# Save results
os.makedirs("logs", exist_ok=True)
with open("logs/solver_judge_countdown.json", "w") as f:
json.dump([episode.to_dict() for episode in results], f, indent=4)
print("\nResults saved to logs/solver_judge_countdown.json")
This script:
- Builds a
TinkerEnginefor rollouts - Wraps it with
AgentWorkflowEngineusingSolverJudgeWorkflow - Executes Countdown tasks and computes pass@1 / pass@k metrics
Monitoring and Checkpoints
- Logging:
- Set
rllm.trainer.logger=['console','wandb']to enable Weights & Biases - Use
rllm.trainer.project_name/rllm.trainer.experiment_nameto organize runs
- Set
- Checkpoints:
- Local paths are controlled by
training.default_local_dir - You can resume from a Tinker checkpoint via
training.resume_from_tinker_id='tinker://<uuid>/weights/<checkpoint_name>'
- Local paths are controlled by
This gives you an end‑to‑end RL training pipeline where rollouts, gradients, and checkpoints all run on Tinker's managed GPU service, while rLLM's unified trainer handles datasets, workflows, advantage computation, and training orchestration.