返回首頁

GPU 與強化學習系統

rl-triton 將七種強化學習回報計算統一為 GPU scan,完整呼叫最高加速 5.70 倍

開源函式庫 rl-triton 把 GAE、V-Trace、Retrace 等信用分配遞迴改寫成 Triton 關聯式掃描,將序列依賴深度由 O(T) 降至 O(log T)。在 4,096 個環境、128 步的測試中,它較已向量化的 torch.compile 基線快 1.6 至 5.70 倍,但大型策略的端到端 PPO 加速僅約 1.02 倍。

Jebulon · CC0 · Image source
zh-Hant

rl-triton 針對強化學習訓練中常被忽略的信用分配階段,發布一組以 Triton 撰寫的融合 GPU kernel。函式庫涵蓋 GAE、V-Trace、Retrace(λ)、TD(λ) returns、discounted returns、eligibility traces 及 episodic prefix sums。作者指出,這七種演算法都可表示成一階線性遞迴 `A_t = α_t + β_t A_{t+1}`,再以同一個結合運算轉成關聯式 scan,將原本逐步執行的 O(T) 依賴鏈改成 O(log T) 平行階段。

演算法專用 kernel 會直接由 rewards、values、done flags 和 importance ratios 在暫存器內建立 α、β,並在晶片上完成中間 scan,避免每個 doubling 階段把結果寫回 HBM。實作亦區分 episode termination、時間截斷與 rollout 視窗邊界,避免 bootstrap 值被錯誤清零;這些細節對 GAE 與 off-policy 方法的數值正確性尤其重要。

作者在 H100 80GB 與 RTX 2000 Ada 上,以 4,096 個平行環境、每段 128 步測試 v0.1.3。相對已經使用 O(log T) scan 的 `torch.compile` 基線,七種運算的完整函式呼叫加速為 1.6 至 5.70 倍;RTX 2000 Ada 的 discounted returns 達最高值,H100 的 GAE 為 2.36 倍。這比拿未編譯逐步迴圈作比較更具參考性。

不過 kernel 微基準不等於整體訓練吞吐。當策略隱藏層為 1024×1024,GAE 只佔 PPO 更新約 2.2%,端到端提升約 1.02 倍;小型 128×128 策略才達 1.11 至 1.16 倍。現版又只接受 FP32,Retrace 在長度 4,096 時會因暫存器壓力慢於編譯基線,兩種 forward scan 超過 131,072 步亦沒有 fallback。工程團隊應先量度信用分配在自身訓練中的比例,再決定是否整合。

來源

  1. rl-triton: High-Performance Triton GPU Kernels for Reinforcement Learning Credit Assignment
  2. simonsays1980/rl-triton