Reinforcement learning kernel optimization speeds training and inference
AReaL-TIK: Stateful Agentic Optimization of Unified RL Kernels through an Optimization IR
Distributed, Parallel, and Cluster Computing
Summary
Training reinforcement learning models on GPUs often uses different code pieces for running simulations and updating the model, which can cause small mismatches affecting performance and results. The authors present KernelBraid, a system that optimizes a single unified code kernel ensuring exact numerical agreement across different training steps. Their method carefully checks correctness while improving speed, achieving about 10% faster throughput on average and up to 2.5 times faster on some tasks. This work helps speed up RL training and inference without changing the learned results.
What this means in practice
- •For machine learning engineers: Improve training efficiency of reinforcement learning models by using a single optimized GPU kernel that ensures exact numerical consistency.
- •For gpu software developers: Develop and optimize GPU kernels for reinforcement learning workloads with a framework that facilitates correctness verification and performance measurement.
Authors
Ran Yan, Youhe Jiang, Jiayi Nie, Wenshuang Li, Yingqi Peng, Taiyi Wang, Tongkai Yang, Binhang Yuan
Abstract
Reinforcement learning (RL) post-training often uses distinct GPU kernels for rollout and policy update. In synchronous PPO and GRPO, numerical disagreement can perturb ratios between current token probabilities and those assigned during rollout. Recomputing rollout log-probabilities with the policy-update backend avoids this discrepancy but adds a forward pass. Bitwise-consistent unified kernels permit reuse when the policy snapshot and probability processing match the objective. Their optimization must preserve agreement across distinct execution regimes. We present KernelBraid, an agentic framework starting from a hand-tuned, bitwise-consistent implementation. Its optimization intermediate representation (IR) organizes source-code search by linking implementations and modifications to numerical requirements, workload measurements, and derivation history. The agent coordinates changes and retains verified intermediates for further exploration; promotion requires passing correctness checks and improving aggregate latency within per-workload limits. Across 12 end-to-end training configurations on H20, KernelBraid achieves 1.10x average throughput relative to AReaL with log-probability recomputation, and the mean training-reward ratio rounds to 1.00x. Isolated-layer profiling yields 1.40x average speedup in summed phase time across 15 model-GPU pairs. Operator-level evaluation covers correctness and performance for 10 operators on A100, H20, and H200, all passing the prescribed bitwise checks. Unified-attention search achieves 2.52x speedup in summed workload latency over the starting implementation using 7M LLM tokens; ablations assess the contributions of retained evidence and branch exploration to search efficiency and attained performance. Our code is open-sourced at https://github.com/areal-project/AReaL-TIK.