🤖 AI Summary
This study addresses the numerical discrepancies and heavy recomputation overhead arising from kernel inconsistencies between rollout and update phases in reinforcement learning post-training. To this end, it proposes KernelBraid, a framework that introduces the first state-preserving, LLM-agent-based approach to kernel optimization. By leveraging intermediate representations (IR) to bridge numerical requirements with derivation histories, the method guides source code search toward the efficient generation of unified GPU kernels with bit-level consistency. Experimental results demonstrate that KernelBraid improves end-to-end training throughput by 1.1×, achieves a 1.4× speedup at the isolated layer level, and reduces attention mechanism latency by 2.52×.
📝 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.