Developing models together, openly
Reinforcement learning (RL) builds decision-making systems that learn from experience to maximize a reward [1][2]. For LLMs it is a key post-training stage that was originally used to create instruction following models [3][4] and more recently has been used to improve performance on verifiable tasks such as math and code [5]. The most prominent open-weight model owes much of its reputation to RL [6], so RL was the natural next step after we pretrained a 32B model in October 2025. At the time, the open-source RL ecosystem for JAX/TPUs was nascent. Scattered work on RL agents existed, but an RL pipeline for LLMs must balance sampling, training, and weight synchronization [7], and none of the existing frameworks handled preemption, which our setting requires. In this post we describe how we built a performant RL pipeline for JAX on TPUs, including the intermediate results, bugs, and missteps along the way.
RL for Marin needed more than sync PPO/GRPO implementation in JAX. We wanted TPU-first execution, asynchronous actor/trainer separation, fast weight sync, high end-to-end throughput, reward/verifier logic for math and code, and checkpointing with restart on preemptible TPU jobs. Our design was constrained by the nature of TPUs on TRC: small preemptible TPU slices and many small inference workers were far easier for us to obtain than one large stable TPU job, so we needed a loose worker-based design.
Tunix — The closest open-source match for LLM RL in JAX on TPUs. It supports PPO/GRPO-style methods, TPU execution, and checkpoint-and-resume. Its async/disaggregated components arrived incrementally through September and October 2025 and were not yet mature in fall 2025. Its disaggregated mode runs as a single tight sub-mesh TPU job rather than as loose workers on small preemptible slices. Its multi-host training requires submitting jobs through Pathways on GKE, which we cannot use.
Brax — The most widely known JAX RL project, with maintained PPO/SAC/ARS training code. It targets physics simulation and classic RL environments, not LLM post-training, and does not provide the trainer/actor/reference/reward/verifier decomposition that LLM RL needs.
RLax — DeepMind’s JAX RL package of reusable primitives. It is not a full system: it provides no rollout system, async trainer/actor architecture, or TPU-native LLM post-training workflow.
PureJaxRL — A compact, fast end-to-end JAX PPO implementation. It is a reference codebase for standard RL environments, not LLM post-training.
Stoix — A JAX RL systems codebase with explicit distributed execution patterns such as Anakin and Sebulba. It remains a single-agent RL research codebase for standard RL environments.
Rejax and EvoRL — JAX RL libraries with PPO support for standard RL training. Neither provides the async LLM rollout, training, and verifier stack we needed.
RLAX (Apple) (paper, related repo: AXLearn) — The closest design to what we wanted: large-scale distributed RL for LLMs on TPUs with trainer/inference separation, verifiers in the loop, and attention to weight sync and preemption. As of March 2026, the paper is being withdrawn, no public RLAX repo exists, and the RL-specific components were never released.
The open JAX RL ecosystem had many PPO implementations but few libraries that addressed TPU-native LLM RL as a systems problem.
We therefore built our own async RL pipeline from scratch. Over five months (November 2025 – March 2026), we went from synthetic baselines with synchronous training to a fully asynchronous system, fixed upstream bugs in open-source libraries along the way, and expanded to harder benchmarks such as AIME and HumanEval+.
Before building anything new, we established baselines using Tinker, a LoRA-based RL system running on GPUs. Thinking Machines found that LoRA matches full fine-tuning for RL [8], so matching Tinker’s results would give us a reasonable baseline for our full fine-tuning pipeline.
Our first milestone was to match Tinker’s results with Marin’s synchronous RL pipeline. Tinker uses an importance-sampling policy-gradient loss that corrects for the mismatch between the policy that sampled a response and the policy being trained. It samples several responses to the same prompt and reinforces the ones that score above the others. We started from a similar objective in Marin, then moved to an RLOO-style loss with leave-one-out advantages.
We began with Llama 3.2 1B. It performed well on synthetic tasks (i.e. three digit addition/multiplication) but poorly on GSM8K, reaching only 0.04 accuracy after 200 steps. Llama 3.1 8B Instruct rose from 0.69 to 0.80 on GSM8K in a single step and from 0.26 to 0.51 on MATH in 180 steps, so we focused on Llama 3.1 8B.
Both Tinker and Marin’s sync RL converged to ~0.43 accuracy on MATH, but Marin took 2x longer (80 steps vs. Tinker’s 35 steps to reach 0.4).
We hypothesize that full fine-tuning disrupted the model’s format-following more than LoRA. Marin’s format accuracy started at 0.47 and took ~80 steps to reach 0.80, so the model spent early training budget re-learning the response format before improving math reasoning. LoRA’s low-rank updates preserve the base model’s capabilities, so Tinker can improve math reasoning from the start (WandB report). A larger sample/train log-probability divergence, since we use vLLM for inference and JAX for training, may have also contributed [9].
Regardless, this was a milestone: We now had reproducible RL training on TPU, confirmed across 3 independent runs.

Tinker (LoRA) vs. Marin (Full FT) on MATH-500. Left: both converge to ~0.43 test accuracy, but Tinker crosses 0.40 at step 29 vs. Marin at step 81 (dashed vertical lines). Center: Marin's format accuracy starts at 0.47 and takes ~80 steps to reach 0.80 (dashed lines), suggesting full fine-tuning disrupted format-following. Right: entropy is similar between both runs, ruling out exploration differences as the cause. (WandB report)

Sync RL runs each stage sequentially. Async RL runs the trainer (Levanter) and actor (vLLM) concurrently with weights synced via Arrow Flight.
Synchronous RL was a simple first step, but each stage (generate, train, eval) completes sequentially, which limits throughput. At this point prior work clearly showed that an async RL system can be performant [7], so that was our next goal.
In December, we built an asynchronous pipeline in which the trainer (Levanter) and actor (vLLM) run concurrently, with model weights synchronized via Arrow Flight. This required two infrastructure changes:
The result: async RL matched sync RL quality (0.26 to 0.50 on MATH-500 in 10 steps) with a 1.21x speedup:
| Metric | Sync RL (wandb) | Async RL (wandb) |
|---|---|---|
| Avg iteration time | 3.71 min | 3.07 min |
| Iterations/minute | 0.269 | 0.326 |
| Median iteration | 3.48 min | 3.02 min |
| Min interval | 3.07 min | 2.40 min |
| Max interval | 5.63 min | 3.82 min |
Unfortunately, we started noticing divergences when moving to async RL. We first noticed two identical async RL runs (i.e. same training config and seed) diverged after dozens of steps. One run peaked at 0.514 accuracy, but the other peaked at 0.482 and then collapsed to 0.36. Confusingly, we found that training metrics (loss, KL, rewards) agreed between the runs, and the divergence appeared only at inference time when we evaluated. (WandB report)

Two identical async RL runs diverge on eval accuracy (left, shaded region) while train accuracy remains indistinguishable (right). Thin lines are raw values and bold lines are EMA-smoothed. The bug only affected sampling at inference time. (WandB report)
We investigated three candidate causes (#2260):
max_tokens=512 left accuracy far above Tinker’s, but did not fix divergence.temp=0.0 instead of 1.0 raised accuracy from 0.294 to 0.442. This was a strong hint, though we did not immediately find the root cause.temp=0 and temp=1 on both platforms finally revealed the bug. On GPU, accuracy dropped from 42.1% to 28.3% as expected. On TPU, it was 40.9% vs. 41.7%: no difference.vLLM on TPU was silently ignoring temperature. All prior async RL evaluations had been effectively greedy.
We traced the bug to input_batch.py in the tpu-inference codebase:
top_k = sampling_params.top_k
if top_k <= 0 or top_k >= vocab_size:
top_k = 1 # BUG: forces greedy!
vLLM’s docs specify that top_k=-1 means “consider all tokens,” but the tpu-inference library converted -1 to 1, selecting only the highest-probability token regardless of temperature! We filed a bug report (tpu-inference #1386) and proposed a fix, which was merged.
This bug also provided a possible explanation for the nondeterminism we observed: We believe that under greedy sampling, small floating-point differences in logit ordering break ties differently across runs.
Separately, we caught a loss normalization regression: switching the DAPO loss from global token normalization to per-example normalization overweighted short responses relative to long reasoning chains and cost 13% accuracy.
After both fixes, MATH-500 accuracy converged to 0.46 (+/-0.02) over 186 steps (WandB run):

After fixing the vLLM top-k bug and loss normalization regression, MATH-500 Pass@1 reaches 0.46 within 10 steps and remains stable (mean=0.45, ±2σ=0.028) over 186 steps of training.
By February, the 186-step run above was the longest we had completed. Our other experiments (Code-R1, AIME) had destabilized around step 240, and no run had yet survived a TPU preemption, so we did not know whether the pipeline could train for longer. In March we migrated the pipeline to Marin’s new Iris scheduler, which gave us an in-cluster coordinator, checkpoint-based resume, and per-phase watchdogs (i.e. a timeout on each phase of a rollout step, so that a hang is reported instead of stalling the run). We then ran three identical 500-step MATH-500 runs (i.e. same config and seed) on Llama 3.1 8B Instruct with RLOO and no KL term (run 1, run 2, run 3).

Three identical 500-step runs (thin: raw, bold: EMA-smoothed). The dashed line marks the end of the previous longest run. Left: held-out MATH-500 Pass@1 peaks at 0.51--0.53 between steps 76 and 247, then drifts down to 0.43--0.45 by step 500. Right: training accuracy peaks at 0.71--0.78 around step 250--360 and also declines. Runs 2 and 3 were preempted twice and once, respectively, and resumed from checkpoint. The resumes are not visible in the curves.
We learned three things from these runs:
Throughput also improved along the way. On the same TPU v5 slice, with the same batch size and the same ~60s forward/backward, median wall-clock per training step dropped from 171s in the January run to 94–103s, and weight-transfer serve time fell from 26s to 8s.
Qwen 2.5 is widely used for post-training, and prior work had shown it to be a stronger base model than Llama for AIME-style math [11], so we wanted it in the pipeline. Supporting it (PR #2446, PR #2456, PR #2458) turned out to require three fixes. First, the model was not registered in tpu-inference, which silently fell back to a slow PyTorch path. Second, the weight sync crashed because Qwen reshapes q_proj differently from Llama. Third, Qwen pads its vocabulary to 152064 tokens for hardware alignment, which conflicted with Levanter’s automatic vocab resizing. With these fixed, we moved to AIME.
MATH-500 had validated the pipeline, but modern models saturate it, so we moved to AIME, the benchmark used by OLMo 3, GLM 4.7, and DeepSeek.
AIME turned out to be hard to evaluate before it was hard to train on. It has only 30 questions, so a single question shifts Pass@1 by 3%, and our first estimates of Pass@k (i.e. the probability that at least one of k samples is correct) were too noisy to read. To reduce this noise we implemented a combinatorial Pass@k estimator (following Codex [12], lighteval, and DeepMath [13]) and increased the eval sample size K per task to 32 (PR #2493).
We then trained Qwen 2.5 7B on DeepMath-103K. Pass@16 improved steadily and reached 0.40, but Pass@1 remained near zero after 40 steps (PR #2441). We hypothesize that Pass@16 must cross some threshold before Pass@1 starts to improve, and that longer training would be needed to reach it.

AIME25 training: Pass@16 steadily improves to 0.40, but Pass@1 remains far from the 0.175 target due to high evaluation variance.
Math was a convenient test bed, but code is the domain with the most practical value, and its verifiers are more complex: a response is correct only if the generated code passes a test suite, so the evaluation environment has to execute that code. Our first code run looked too good to be true. Accuracy climbed to ~100% within 26 steps, and when we looked closer we found that the evaluation environment executed the test scripts without ever invoking the validation function.
After fixing the eval, we reproduced Code-R1’s results [10] by training Qwen 2.5 7B Instruct with RL on 2K LeetCode questions (PR #2286). HumanEval+ improved from 0.80 to 0.84 in 264 steps, matching Code-R1’s reported 0.848 (wandb run). Pass@1 then destabilized after 240 steps. We believe this is because we omitted the KL term that Code-R1 uses [10].

Left: bugged verifier falsely showed ~100% accuracy. Right: after fixing the eval, HumanEval+ Pass@1 improves from 0.80 to 0.84, closely matching Code-R1's reported 0.848 (dashed line). Pass@1 destabilizes after ~240 steps.
At this point we are shifting from RL to SFT for the next Marin model release. Three things are at the top of the list when we return to RL:
We gratefully acknowledge Google’s TPU Research Cloud (TRC) program for providing the TPU resources that made this work possible.
[1] Sutton, R.S. and Barto, A.G. (2018). Reinforcement Learning: An Introduction, 2nd Edition. MIT Press.
[2] Silver, D., Huang, A., Maddison, C. et al. (2016). Mastering the game of Go with deep neural networks and tree search. Nature, 529(7587), 484-489.
[3] Bai, Y., Kadavath, S., Kundu, S. et al. (2022). Constitutional AI: Harmlessness from AI Feedback. arXiv:2212.08073.
[4] Ouyang, L., Wu, J., Jiang, X. et al. (2022). Training language models to follow instructions with human feedback. NeurIPS 2022. arXiv:2203.02155.
[5] OpenAI (2024). OpenAI o1 System Card. arXiv:2412.16720.
[6] DeepSeek-AI et al. (2025). DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via Reinforcement Learning. arXiv:2501.12948.
[7] Mistral AI et al. (2025). Magistral. arXiv:2506.10910.
[8] Thinking Machines Lab (2025). LoRA Without Regret.
[9] Zheng, C. et al. (2025). Defeating the Training-Inference Mismatch via FP16. arXiv:2510.26788.
[10] Liu, J. et al. (2025). Code-R1: Reproducing R1 for Code with Reliable Rewards.
[11] Liu, Z., Chen, Z., Li, J. et al. (2025). Understanding R1-Zero-Like Training: A Critical Perspective (Dr. GRPO). COLM 2025. arXiv:2503.20783.
[12] Chen, M., Tworek, J., Jun, H. et al. (2021). Evaluating Large Language Models Trained on Code. arXiv:2107.03374.
[13] He, Z. et al. (2025). DeepMath-103K: A Large-Scale, Challenging, Decontaminated, and Verifiable Mathematical Dataset for Advancing Reasoning. arXiv:2504.11456.