55 KiB
PPO CUDA Pipeline Phase 2c — Implementation Plan
For Claude: REQUIRED SUB-SKILL: Use superpowers:executing-plans to implement this plan task-by-task.
Goal: Build a zero-roundtrip CUDA kernel for PPO experience collection — actor forward, softmax sampling, critic forward, portfolio simulation, barrier tracking, diversity entropy, curiosity inference, reward combination, and GAE advantage computation — with 128 parallel episodes × 500 timesteps = 64,000 training-ready experience tuples per kernel launch.
Architecture: Single monolithic kernel ppo_full_experience_kernel in ppo_experience_kernel.cu. Shared device functions extracted from DQN kernel into common_device_functions.cuh and reused by both kernels via NVRTC source concatenation. Thread mapping: 1 thread = 1 episode. Two-phase: forward rollout (collect experiences) then backward GAE scan (compute advantages and returns).
Tech Stack: CUDA (NVRTC runtime compilation), cudarc 0.17.3 via candle_core::cuda_backend::cudarc, Rust #[cfg(feature = "cuda")]
Design Doc: docs/plans/2026-02-28-ppo-cuda-pipeline-design.md
Task 1: Extract shared device functions to common_device_functions.cuh
Files:
- Create:
crates/ml/src/cuda_pipeline/common_device_functions.cuh - Modify:
crates/ml/src/cuda_pipeline/dqn_experience_kernel.cu
Context: Lines 1-382 of dqn_experience_kernel.cu contain device helper functions that both DQN and PPO kernels need. Extract them into a shared header. The DQN kernel currently has these as inline definitions. After extraction, DQN will include them via NVRTC source concatenation.
Step 1: Create common_device_functions.cuh
Copy the following functions from dqn_experience_kernel.cu into the new header file:
// common_device_functions.cuh — Shared device functions for DQN/PPO CUDA kernels
//
// Included via NVRTC source concatenation (not #include).
// Both kernels prepend this file's contents before their own source.
/* Constants shared by all experience collection kernels */
#define STATE_DIM 54
#define MARKET_DIM 51
#define PORTFOLIO_DIM 3
#define NUM_ACTIONS 45
/* Curiosity forward model */
#define CUR_INPUT 35
#define CUR_HIDDEN 64
#define CUR_OUTPUT 32
/* Auxiliary state */
#define DIVERSITY_WINDOW 100
#define PORTFOLIO_STATE_SIZE 8
#define BARRIER_STATE_SIZE 5
Then copy these device functions verbatim from dqn_experience_kernel.cu:
gpu_random(line 42-45)leaky_relu(line 48-50)matvec_leaky_relu(line 58-75)action_to_exposure(line 166-176)action_to_tx_cost(line 184-192)barrier_init(line 200-215)barrier_check(line 223-259)barrier_reset(line 262-268)diversity_entropy(line 280-326)curiosity_inference(line 338-382)
Do NOT copy q_forward_dueling or argmax_q — those are DQN-specific.
Step 2: Remove extracted functions from dqn_experience_kernel.cu
Remove lines 1-382 (everything from the file header through curiosity_inference). Replace with:
/**
* Zero-Roundtrip DQN Experience Collection Kernel
*
* Requires common_device_functions.cuh prepended via NVRTC source concatenation.
* Launch config: grid=(ceil(N/32),1,1), block=(32,1,1).
* Each thread processes one independent episode of L timesteps.
* 128 episodes x 500 timesteps = 64,000 experiences per launch.
*/
/* DQN-specific layer sizes */
#define SHARED_H1 256
#define SHARED_H2 256
#define VALUE_H 128
#define ADV_H 128
Then keep q_forward_dueling and argmax_q functions (moved to top of file after the defines), followed by the main kernel entry point unchanged.
Step 3: Run test to verify DQN kernel source is still valid
Run: SQLX_OFFLINE=true cargo test -p ml --lib test_kernel_source_contains_entry_point -- --exact
Expected: PASS (the test checks dqn_experience_kernel.cu for dqn_full_experience_kernel and extern "C")
Step 4: Update gpu_experience_collector.rs for source concatenation
In crates/ml/src/cuda_pipeline/gpu_experience_collector.rs, change the kernel compilation (around line 194) from:
let kernel_src = include_str!("dqn_experience_kernel.cu");
let ptx: Ptx = cudarc::nvrtc::compile_ptx(kernel_src)
to:
let common_src = include_str!("common_device_functions.cuh");
let kernel_src = include_str!("dqn_experience_kernel.cu");
let full_source = format!("{}\n{}", common_src, kernel_src);
let ptx: Ptx = cudarc::nvrtc::compile_ptx(&full_source)
Step 5: Add test for common header
In crates/ml/src/cuda_pipeline/mod.rs, add:
#[test]
fn test_common_header_contains_shared_functions() {
let src = include_str!("common_device_functions.cuh");
assert!(src.contains("gpu_random"), "Missing gpu_random");
assert!(src.contains("matvec_leaky_relu"), "Missing matvec_leaky_relu");
assert!(src.contains("action_to_exposure"), "Missing action_to_exposure");
assert!(src.contains("barrier_init"), "Missing barrier_init");
assert!(src.contains("diversity_entropy"), "Missing diversity_entropy");
assert!(src.contains("curiosity_inference"), "Missing curiosity_inference");
// Must NOT contain DQN-specific functions
assert!(!src.contains("q_forward_dueling"), "DQN-specific function leaked into common header");
assert!(!src.contains("dqn_full_experience_kernel"), "DQN kernel leaked into common header");
}
Step 6: Run all cuda_pipeline tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib cuda_pipeline -- --test-threads=1
Expected: 15/15 PASS (14 existing + 1 new)
Step 7: Commit
git add crates/ml/src/cuda_pipeline/common_device_functions.cuh \
crates/ml/src/cuda_pipeline/dqn_experience_kernel.cu \
crates/ml/src/cuda_pipeline/gpu_experience_collector.rs \
crates/ml/src/cuda_pipeline/mod.rs
git commit -m "refactor(cuda): extract shared device functions to common_device_functions.cuh"
Task 2: PPO actor forward + softmax sampling device functions
Files:
- Create:
crates/ml/src/cuda_pipeline/ppo_experience_kernel.cu
Context: This creates the PPO kernel file with actor-specific device functions. The shared functions (RNG, matmul, portfolio, barriers, diversity, curiosity) come from common_device_functions.cuh via source concatenation — do NOT redefine them.
Step 1: Create ppo_experience_kernel.cu with PPO constants and actor functions
/**
* Zero-Roundtrip PPO Experience Collection Kernel
*
* Requires common_device_functions.cuh prepended via NVRTC source concatenation.
* Launch config: grid=(ceil(N/32),1,1), block=(32,1,1).
* Each thread runs one independent episode.
*
* Phase A: Forward rollout (L timesteps) — actor, critic, portfolio, rewards
* Phase B: Backward GAE scan — advantages and returns
*/
/* PPO Actor layer sizes: 54 -> 128 (ReLU) -> 64 (ReLU) -> 45 */
#define ACTOR_H1 128
#define ACTOR_H2 64
/* PPO Critic layer sizes: 54 -> 512 -> 384 -> 256 -> 128 -> 64 -> 1 (all ReLU) */
#define CRITIC_H1 512
#define CRITIC_H2 384
#define CRITIC_H3 256
#define CRITIC_H4 128
#define CRITIC_H5 64
/* GAE hyperparameters (passed as kernel args for flexibility) */
/* #define GAE_GAMMA and GAE_LAMBDA are NOT used — passed as args */
/* ------------------------------------------------------------------ */
/* PPO Actor Forward */
/* ------------------------------------------------------------------ */
/**
* PPO actor forward pass: state[54] -> h1[128] (ReLU) -> h2[64] (ReLU) -> logits[45]
*
* Uses matvec_leaky_relu with activate=1 (LeakyReLU with alpha=0.01 serves as ReLU
* for non-negative inputs; the 0.01 leak is negligible for positive activations).
* Final layer has no activation (raw logits for softmax).
*/
__device__ void ppo_actor_forward(
const float* state, /* [STATE_DIM] */
const float* __restrict__ pw1, /* [ACTOR_H1, STATE_DIM] */
const float* __restrict__ pb1, /* [ACTOR_H1] */
const float* __restrict__ pw2, /* [ACTOR_H2, ACTOR_H1] */
const float* __restrict__ pb2, /* [ACTOR_H2] */
const float* __restrict__ pw3, /* [NUM_ACTIONS, ACTOR_H2] */
const float* __restrict__ pb3, /* [NUM_ACTIONS] */
float* h1, /* scratch [ACTOR_H1] */
float* h2, /* scratch [ACTOR_H2] */
float* logits /* output [NUM_ACTIONS] */
) {
matvec_leaky_relu(pw1, pb1, state, h1, STATE_DIM, ACTOR_H1, 1);
matvec_leaky_relu(pw2, pb2, h1, h2, ACTOR_H1, ACTOR_H2, 1);
matvec_leaky_relu(pw3, pb3, h2, logits, ACTOR_H2, NUM_ACTIONS, 0); /* no activation */
}
/**
* Stable softmax + categorical sampling.
*
* 1. Subtract max for numerical stability
* 2. Exponentiate and sum
* 3. Normalize to probabilities
* 4. CDF scan + LCG random draw → action index
* 5. Compute log(p[action]) for PPO loss
*
* Returns (action_idx, log_prob) via out parameters.
*/
__device__ void softmax_sample(
const float* logits, /* [NUM_ACTIONS] */
float* probs, /* scratch+output [NUM_ACTIONS] */
unsigned int* rng,
int* out_action,
float* out_log_prob
) {
/* Step 1: Find max for stability */
float max_logit = logits[0];
for (int i = 1; i < NUM_ACTIONS; i++) {
if (logits[i] > max_logit) max_logit = logits[i];
}
/* Step 2: exp(logit - max) and sum */
float sum_exp = 0.0f;
for (int i = 0; i < NUM_ACTIONS; i++) {
probs[i] = expf(logits[i] - max_logit);
sum_exp += probs[i];
}
/* Step 3: Normalize */
float inv_sum = 1.0f / (sum_exp + 1e-8f);
for (int i = 0; i < NUM_ACTIONS; i++) {
probs[i] *= inv_sum;
}
/* Step 4: CDF scan + random draw */
float u = gpu_random(rng);
float cdf = 0.0f;
int action = NUM_ACTIONS - 1; /* fallback to last action */
for (int i = 0; i < NUM_ACTIONS; i++) {
cdf += probs[i];
if (u < cdf) {
action = i;
break;
}
}
/* Step 5: log probability */
*out_action = action;
*out_log_prob = logf(probs[action] + 1e-8f);
}
Step 2: Add source verification test in mod.rs
#[test]
fn test_ppo_kernel_source_contains_actor_functions() {
let src = include_str!("ppo_experience_kernel.cu");
assert!(src.contains("ppo_actor_forward"), "Missing ppo_actor_forward");
assert!(src.contains("softmax_sample"), "Missing softmax_sample");
}
Step 3: Run tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib cuda_pipeline -- --test-threads=1
Expected: 16/16 PASS
Step 4: Commit
git add crates/ml/src/cuda_pipeline/ppo_experience_kernel.cu \
crates/ml/src/cuda_pipeline/mod.rs
git commit -m "feat(cuda): add PPO actor forward and softmax sampling device functions"
Task 3: PPO critic forward + GAE backward scan device functions
Files:
- Modify:
crates/ml/src/cuda_pipeline/ppo_experience_kernel.cu
Context: Add the critic forward pass (5-layer deep MLP with ping-pong scratch buffers) and GAE backward scan to the PPO kernel file.
Step 1: Add critic forward function
Append after softmax_sample:
/* ------------------------------------------------------------------ */
/* PPO Critic Forward */
/* ------------------------------------------------------------------ */
/**
* PPO critic forward pass (5-layer deep):
* state[54] -> 512 (ReLU) -> 384 (ReLU) -> 256 (ReLU) -> 128 (ReLU) -> 64 (ReLU) -> 1
*
* Uses ping-pong pattern: two 512-wide scratch buffers alternate to avoid
* allocating separate arrays for each layer. Layer outputs always fit
* within 512 floats (the largest hidden dim).
*/
__device__ float ppo_critic_forward(
const float* state, /* [STATE_DIM] */
const float* __restrict__ vw1, /* [CRITIC_H1, STATE_DIM] */
const float* __restrict__ vb1, /* [CRITIC_H1] */
const float* __restrict__ vw2, /* [CRITIC_H2, CRITIC_H1] */
const float* __restrict__ vb2, /* [CRITIC_H2] */
const float* __restrict__ vw3, /* [CRITIC_H3, CRITIC_H2] */
const float* __restrict__ vb3, /* [CRITIC_H3] */
const float* __restrict__ vw4, /* [CRITIC_H4, CRITIC_H3] */
const float* __restrict__ vb4, /* [CRITIC_H4] */
const float* __restrict__ vw5, /* [CRITIC_H5, CRITIC_H4] */
const float* __restrict__ vb5, /* [CRITIC_H5] */
const float* __restrict__ vw6, /* [1, CRITIC_H5] */
const float* __restrict__ vb6, /* [1] */
float* scratch_a, /* [CRITIC_H1] (ping) */
float* scratch_b /* [CRITIC_H1] (pong) */
) {
/* Layer 1: state[54] -> scratch_a[512] */
matvec_leaky_relu(vw1, vb1, state, scratch_a, STATE_DIM, CRITIC_H1, 1);
/* Layer 2: scratch_a[512] -> scratch_b[384] */
matvec_leaky_relu(vw2, vb2, scratch_a, scratch_b, CRITIC_H1, CRITIC_H2, 1);
/* Layer 3: scratch_b[384] -> scratch_a[256] */
matvec_leaky_relu(vw3, vb3, scratch_b, scratch_a, CRITIC_H2, CRITIC_H3, 1);
/* Layer 4: scratch_a[256] -> scratch_b[128] */
matvec_leaky_relu(vw4, vb4, scratch_a, scratch_b, CRITIC_H3, CRITIC_H4, 1);
/* Layer 5: scratch_b[128] -> scratch_a[64] */
matvec_leaky_relu(vw5, vb5, scratch_b, scratch_a, CRITIC_H4, CRITIC_H5, 1);
/* Output: scratch_a[64] -> scalar value */
float value = vb6[0];
for (int i = 0; i < CRITIC_H5; i++) {
value += vw6[i] * scratch_a[i];
}
return value;
}
/* ------------------------------------------------------------------ */
/* GAE Backward Scan */
/* ------------------------------------------------------------------ */
/**
* Generalized Advantage Estimation — backward scan over one episode.
*
* For each timestep t (scanning backward from L-1 to 0):
* delta[t] = rewards[t] + gamma * values[t+1] * (1 - dones[t]) - values[t]
* advantages[t] = delta[t] + gamma * lambda * (1 - dones[t]) * advantages[t+1]
* returns[t] = advantages[t] + values[t]
*
* values[L] is treated as the bootstrap value (value of last next-state).
* If the last step is done, bootstrap is 0.
*/
__device__ void compute_gae_backward(
const float* rewards, /* [L] */
const float* values, /* [L+1] — values[L] is bootstrap */
const int* dones, /* [L] */
float* advantages, /* [L] output */
float* returns, /* [L] output */
int L,
float gamma,
float lambda
) {
float gae = 0.0f;
for (int t = L - 1; t >= 0; t--) {
float not_done = (dones[t] == 0) ? 1.0f : 0.0f;
float delta = rewards[t] + gamma * values[t + 1] * not_done - values[t];
gae = delta + gamma * lambda * not_done * gae;
advantages[t] = gae;
returns[t] = gae + values[t];
}
}
Step 2: Add source verification test
In crates/ml/src/cuda_pipeline/mod.rs, update the existing PPO source test:
#[test]
fn test_ppo_kernel_source_contains_all_functions() {
let src = include_str!("ppo_experience_kernel.cu");
assert!(src.contains("ppo_actor_forward"), "Missing ppo_actor_forward");
assert!(src.contains("softmax_sample"), "Missing softmax_sample");
assert!(src.contains("ppo_critic_forward"), "Missing ppo_critic_forward");
assert!(src.contains("compute_gae_backward"), "Missing compute_gae_backward");
}
(Replace the Task 2 test if it already exists, since this is a superset.)
Step 3: Run tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib cuda_pipeline -- --test-threads=1
Expected: 16/16 PASS
Step 4: Commit
git add crates/ml/src/cuda_pipeline/ppo_experience_kernel.cu \
crates/ml/src/cuda_pipeline/mod.rs
git commit -m "feat(cuda): add PPO critic forward (5-layer ping-pong) and GAE backward scan"
Task 4: PPO main kernel entry point
Files:
- Modify:
crates/ml/src/cuda_pipeline/ppo_experience_kernel.cu
Context: Add the main ppo_full_experience_kernel entry point that wires all device functions together. Phase A: forward rollout (L timesteps per thread). Phase B: backward GAE scan. Pattern mirrors dqn_full_experience_kernel (lines 394-785 of original DQN kernel) but replaces Q-network + epsilon-greedy with actor softmax + critic value estimation, and adds GAE as a second pass.
Step 1: Add the main kernel
Append to ppo_experience_kernel.cu:
/* ------------------------------------------------------------------ */
/* Main Kernel */
/* ------------------------------------------------------------------ */
/**
* Full PPO experience collection kernel.
*
* Phase A: Forward rollout — actor, critic, portfolio, rewards, store per-step data
* Phase B: Backward GAE — compute advantages and returns from stored values/rewards
*
* Each thread runs one independent episode of L timesteps.
* Grid: (ceil(N/32), 1, 1), Block: (32, 1, 1).
*/
extern "C" __global__ void ppo_full_experience_kernel(
/* Market data [total_bars, MARKET_DIM] */
const float* __restrict__ market_features,
/* Target prices [total_bars, 4]: preproc_close, preproc_next, raw_close, raw_next */
const float* __restrict__ targets,
/* Episode start indices [N] */
const int* __restrict__ episode_starts,
/* ---- Actor weights (6 pointers) ---- */
const float* __restrict__ pw1, /* [ACTOR_H1, STATE_DIM] */
const float* __restrict__ pb1, /* [ACTOR_H1] */
const float* __restrict__ pw2, /* [ACTOR_H2, ACTOR_H1] */
const float* __restrict__ pb2, /* [ACTOR_H2] */
const float* __restrict__ pw3, /* [NUM_ACTIONS, ACTOR_H2] */
const float* __restrict__ pb3, /* [NUM_ACTIONS] */
/* ---- Critic weights (12 pointers) ---- */
const float* __restrict__ vw1, /* [CRITIC_H1, STATE_DIM] */
const float* __restrict__ vb1, /* [CRITIC_H1] */
const float* __restrict__ vw2, /* [CRITIC_H2, CRITIC_H1] */
const float* __restrict__ vb2, /* [CRITIC_H2] */
const float* __restrict__ vw3, /* [CRITIC_H3, CRITIC_H2] */
const float* __restrict__ vb3, /* [CRITIC_H3] */
const float* __restrict__ vw4, /* [CRITIC_H4, CRITIC_H3] */
const float* __restrict__ vb4, /* [CRITIC_H4] */
const float* __restrict__ vw5, /* [CRITIC_H5, CRITIC_H4] */
const float* __restrict__ vb5, /* [CRITIC_H5] */
const float* __restrict__ vw6, /* [1, CRITIC_H5] */
const float* __restrict__ vb6, /* [1] */
/* ---- Curiosity weights (4 pointers) ---- */
const float* __restrict__ cw1, /* [CUR_HIDDEN, CUR_INPUT] */
const float* __restrict__ cb1, /* [CUR_HIDDEN] */
const float* __restrict__ cw2, /* [CUR_OUTPUT, CUR_HIDDEN] */
const float* __restrict__ cb2, /* [CUR_OUTPUT] */
/* ---- Per-episode mutable state arrays ---- */
float* portfolio_states, /* [N, PORTFOLIO_STATE_SIZE] */
float* barrier_states, /* [N, BARRIER_STATE_SIZE] */
int* diversity_windows, /* [N, DIVERSITY_WINDOW] */
int* diversity_metas, /* [N, 2] */
/* ---- Barrier config (shared) ---- */
const float* __restrict__ barrier_config, /* [3]: profit_mult, loss_mult, max_bars */
/* ---- Scalar configs ---- */
float max_position,
int episode_length,
int total_bars,
int L, /* timesteps per episode */
float gamma, /* discount factor for GAE */
float gae_lambda, /* lambda for GAE */
float curiosity_max_reward,
int N, /* total number of episodes */
float barrier_scale,
float diversity_scale,
float curiosity_scale,
float risk_weight,
/* ---- RNG states [N] ---- */
unsigned int* rng_states,
/* ---- Output arrays ---- */
float* out_states, /* [N, L, STATE_DIM] */
int* out_actions, /* [N, L] */
float* out_log_probs, /* [N, L] */
float* out_advantages, /* [N, L] */
float* out_returns, /* [N, L] */
int* out_dones /* [N, L] */
) {
int tid = blockIdx.x * blockDim.x + threadIdx.x;
if (tid >= N) return;
/* ---- Load per-thread portfolio state ---- */
int ps_off = tid * PORTFOLIO_STATE_SIZE;
float cash = portfolio_states[ps_off + 0];
float position = portfolio_states[ps_off + 1];
float entry_price = portfolio_states[ps_off + 2];
float initial_cap = portfolio_states[ps_off + 3];
float spread = portfolio_states[ps_off + 4];
float last_price = portfolio_states[ps_off + 5];
float reserve_pct = portfolio_states[ps_off + 6];
float cum_costs = portfolio_states[ps_off + 7];
int bs_off = tid * BARRIER_STATE_SIZE;
float barrier_st[BARRIER_STATE_SIZE];
for (int i = 0; i < BARRIER_STATE_SIZE; i++)
barrier_st[i] = barrier_states[bs_off + i];
int dw_off = tid * DIVERSITY_WINDOW;
int div_window[DIVERSITY_WINDOW];
for (int i = 0; i < DIVERSITY_WINDOW; i++)
div_window[i] = diversity_windows[dw_off + i];
int dm_off = tid * 2;
int div_meta[2];
div_meta[0] = diversity_metas[dm_off + 0];
div_meta[1] = diversity_metas[dm_off + 1];
unsigned int rng = rng_states[tid];
int global_bar = episode_starts[tid];
/* ---- Per-thread scratch buffers ---- */
float state[STATE_DIM];
float actor_h1[ACTOR_H1];
float actor_h2[ACTOR_H2];
float logits[NUM_ACTIONS];
float probs[NUM_ACTIONS];
float critic_a[CRITIC_H1]; /* ping-pong buffer A */
float critic_b[CRITIC_H1]; /* ping-pong buffer B */
float next_state[STATE_DIM];
float cur_scratch[CUR_HIDDEN]; /* curiosity scratch */
/* ---- GAE storage arrays (per-thread, per-timestep) ---- */
float gae_values[500 + 1]; /* values[0..L] + bootstrap at [L] */
float gae_rewards[500];
int gae_dones[500];
int step_in_episode = 0;
/* ============================================================== */
/* Phase A: Forward Rollout */
/* ============================================================== */
for (int t = 0; t < L; t++) {
int bar = global_bar % total_bars;
/* 1. Read market features */
int feat_off = bar * MARKET_DIM;
for (int i = 0; i < MARKET_DIM; i++)
state[i] = market_features[feat_off + i];
/* 2. Compute portfolio features */
float portfolio_value = cash + position * last_price;
float norm_value = (initial_cap > 0.0f) ? (portfolio_value / initial_cap) : 1.0f;
state[MARKET_DIM + 0] = norm_value;
state[MARKET_DIM + 1] = position / (max_position > 0.0f ? max_position : 1.0f);
state[MARKET_DIM + 2] = spread;
/* 3. Actor forward pass */
ppo_actor_forward(state, pw1, pb1, pw2, pb2, pw3, pb3,
actor_h1, actor_h2, logits);
/* 4. Softmax + categorical sample */
int action;
float log_prob;
softmax_sample(logits, probs, &rng, &action, &log_prob);
/* 5. Critic forward pass */
float value = ppo_critic_forward(state, vw1, vb1, vw2, vb2, vw3, vb3,
vw4, vb4, vw5, vb5, vw6, vb6,
critic_a, critic_b);
/* 6. Portfolio simulation (reuse from DQN) */
float target_exposure = action_to_exposure(action);
float tx_cost_rate = action_to_tx_cost(action);
int tgt_off = bar * 4;
float raw_close = targets[tgt_off + 2];
float raw_next = targets[tgt_off + 3];
if (raw_close <= 0.0f) raw_close = 1.0f;
if (raw_next <= 0.0f) raw_next = raw_close;
float old_position = position;
float desired_pos = target_exposure * max_position;
float delta_pos = desired_pos - position;
float tx_cost = 0.0f;
if (delta_pos != 0.0f) {
float notional = fabsf(delta_pos) * raw_close;
tx_cost = notional * tx_cost_rate;
cash -= tx_cost;
cum_costs += tx_cost;
float pos_cost = delta_pos * raw_close;
cash -= pos_cost;
position = desired_pos;
/* Barrier: open on new position */
if (fabsf(old_position) < 0.01f && fabsf(position) > 0.01f) {
barrier_init(barrier_st, barrier_config, raw_close, step_in_episode);
}
}
/* Update portfolio with price change */
float pnl = position * (raw_next - raw_close);
cash += pnl;
last_price = raw_next;
entry_price = (fabsf(position) > 0.01f) ? entry_price : 0.0f;
if (entry_price <= 0.0f && fabsf(position) > 0.01f) {
entry_price = raw_close;
}
/* 7. Barrier tracking */
int barrier_label = barrier_check(barrier_st, raw_next, step_in_episode, position);
float barrier_reward = 0.0f;
if (barrier_label != 0) {
barrier_reward = (float)barrier_label;
barrier_reset(barrier_st);
}
/* 8. Diversity entropy */
float div_penalty = diversity_entropy(div_window, div_meta, action);
/* 9. Curiosity inference */
/* Read next_state features for curiosity target */
int next_bar = (bar + 1) % total_bars;
int next_off = next_bar * MARKET_DIM;
for (int i = 0; i < MARKET_DIM; i++)
next_state[i] = market_features[next_off + i];
/* Fill portfolio dims in next_state */
float next_portfolio_value = cash + position * raw_next;
float next_norm_value = (initial_cap > 0.0f) ? (next_portfolio_value / initial_cap) : 1.0f;
next_state[MARKET_DIM + 0] = next_norm_value;
next_state[MARKET_DIM + 1] = position / (max_position > 0.0f ? max_position : 1.0f);
next_state[MARKET_DIM + 2] = spread;
float curiosity_bonus = curiosity_inference(
state, next_state, action,
cw1, cb1, cw2, cb2, cur_scratch,
curiosity_max_reward
);
/* 10. Risk penalty */
float drawdown = 1.0f - (norm_value > 0.0f ? norm_value : 1.0f);
if (drawdown < 0.0f) drawdown = 0.0f;
float risk_penalty = risk_weight * drawdown;
/* 11. Reward combination */
float log_return = (raw_close > 0.0f) ? logf(raw_next / raw_close) : 0.0f;
float base_reward = log_return * 1000.0f;
if (fabsf(position) > 0.01f) {
base_reward *= (position > 0.0f) ? 1.0f : -1.0f;
} else {
base_reward = 0.0f;
}
float reward = base_reward * (1.0f + barrier_scale * barrier_reward)
+ diversity_scale * div_penalty
+ curiosity_scale * curiosity_bonus
- risk_penalty;
/* Episode termination */
step_in_episode++;
int done = (step_in_episode >= episode_length) ? 1 : 0;
/* 12. Store per-timestep outputs */
int out_off = tid * L + t;
int state_off = out_off * STATE_DIM;
for (int i = 0; i < STATE_DIM; i++)
out_states[state_off + i] = state[i];
out_actions[out_off] = action;
out_log_probs[out_off] = log_prob;
out_dones[out_off] = done;
/* Store for GAE */
gae_values[t] = value;
gae_rewards[t] = reward;
gae_dones[t] = done;
/* Episode reset on done */
if (done) {
step_in_episode = 0;
/* Reset portfolio */
cash = initial_cap;
position = 0.0f;
entry_price = 0.0f;
last_price = 0.0f;
cum_costs = 0.0f;
barrier_reset(barrier_st);
/* Reset diversity window */
for (int i = 0; i < DIVERSITY_WINDOW; i++) div_window[i] = 0;
div_meta[0] = 0;
div_meta[1] = 0;
}
global_bar++;
}
/* ============================================================== */
/* Phase B: Backward GAE Scan */
/* ============================================================== */
/* Bootstrap value: if last step was done, bootstrap=0; else run critic on next_state */
if (gae_dones[L - 1]) {
gae_values[L] = 0.0f;
} else {
/* Need to build the "next state" after the last timestep and run critic */
int last_bar = (global_bar - 1) % total_bars;
int bootstrap_bar = (last_bar + 1) % total_bars;
int b_off = bootstrap_bar * MARKET_DIM;
float bootstrap_state[STATE_DIM];
for (int i = 0; i < MARKET_DIM; i++)
bootstrap_state[i] = market_features[b_off + i];
/* Approximate portfolio features (use current values) */
float boot_pv = cash + position * last_price;
float boot_nv = (initial_cap > 0.0f) ? (boot_pv / initial_cap) : 1.0f;
bootstrap_state[MARKET_DIM + 0] = boot_nv;
bootstrap_state[MARKET_DIM + 1] = position / (max_position > 0.0f ? max_position : 1.0f);
bootstrap_state[MARKET_DIM + 2] = spread;
gae_values[L] = ppo_critic_forward(bootstrap_state,
vw1, vb1, vw2, vb2, vw3, vb3,
vw4, vb4, vw5, vb5, vw6, vb6,
critic_a, critic_b);
}
/* Run GAE backward scan */
float gae_advantages[500];
float gae_returns[500];
compute_gae_backward(gae_rewards, gae_values, gae_dones,
gae_advantages, gae_returns, L, gamma, gae_lambda);
/* Write advantages and returns to output */
int base = tid * L;
for (int t = 0; t < L; t++) {
out_advantages[base + t] = gae_advantages[t];
out_returns[base + t] = gae_returns[t];
}
/* ---- Write back per-thread state ---- */
portfolio_states[ps_off + 0] = cash;
portfolio_states[ps_off + 1] = position;
portfolio_states[ps_off + 2] = entry_price;
portfolio_states[ps_off + 3] = initial_cap;
portfolio_states[ps_off + 4] = spread;
portfolio_states[ps_off + 5] = last_price;
portfolio_states[ps_off + 6] = reserve_pct;
portfolio_states[ps_off + 7] = cum_costs;
for (int i = 0; i < BARRIER_STATE_SIZE; i++)
barrier_states[bs_off + i] = barrier_st[i];
for (int i = 0; i < DIVERSITY_WINDOW; i++)
diversity_windows[dw_off + i] = div_window[i];
diversity_metas[dm_off + 0] = div_meta[0];
diversity_metas[dm_off + 1] = div_meta[1];
rng_states[tid] = rng;
}
Step 2: Update source test
Ensure the test from Task 2/3 now also checks:
assert!(src.contains("ppo_full_experience_kernel"), "Missing kernel entry point");
assert!(src.contains("extern \"C\""), "Missing extern C linkage");
Step 3: Run tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib cuda_pipeline -- --test-threads=1
Expected: 16/16 PASS
Step 4: Commit
git add crates/ml/src/cuda_pipeline/ppo_experience_kernel.cu \
crates/ml/src/cuda_pipeline/mod.rs
git commit -m "feat(cuda): add PPO main kernel — forward rollout + backward GAE scan"
Task 5: PPO weight extraction (PpoActorWeightSet, PpoCriticWeightSet)
Files:
- Modify:
crates/ml/src/cuda_pipeline/gpu_weights.rs
Context: Add PPO actor (6 tensors) and critic (12 tensors) weight sets following the existing DuelingWeightSet/CuriosityWeightSet pattern. VarMap key names come from crates/ml/src/ppo/ppo.rs: policy_layer_0.weight, policy_layer_0.bias, etc.
Step 1: Add weight name constants
After CURIOSITY_WEIGHT_NAMES:
/// PPO Actor weight tensor names (6 total).
const PPO_ACTOR_WEIGHT_NAMES: [&str; 6] = [
"policy_layer_0.weight",
"policy_layer_0.bias",
"policy_layer_1.weight",
"policy_layer_1.bias",
"policy_output.weight",
"policy_output.bias",
];
/// PPO Critic weight tensor names (12 total).
const PPO_CRITIC_WEIGHT_NAMES: [&str; 12] = [
"value_layer_0.weight",
"value_layer_0.bias",
"value_layer_1.weight",
"value_layer_1.bias",
"value_layer_2.weight",
"value_layer_2.bias",
"value_layer_3.weight",
"value_layer_3.bias",
"value_layer_4.weight",
"value_layer_4.bias",
"value_output.weight",
"value_output.bias",
];
Step 2: Add structs
After CuriosityWeightSet:
/// GPU buffers holding all 6 weight tensors of the PPO Actor (PolicyNetwork).
///
/// Layout matches `PolicyNetwork` in `crates/ml/src/ppo/ppo.rs`:
/// - Layer 0: `[128, 54]`, `[128]`
/// - Layer 1: `[64, 128]`, `[64]`
/// - Output: `[45, 64]`, `[45]`
#[allow(missing_debug_implementations)]
pub struct PpoActorWeightSet {
pub pw1: CudaSlice<f32>, // policy_layer_0.weight [128, 54]
pub pb1: CudaSlice<f32>, // policy_layer_0.bias [128]
pub pw2: CudaSlice<f32>, // policy_layer_1.weight [64, 128]
pub pb2: CudaSlice<f32>, // policy_layer_1.bias [64]
pub pw3: CudaSlice<f32>, // policy_output.weight [45, 64]
pub pb3: CudaSlice<f32>, // policy_output.bias [45]
}
/// GPU buffers holding all 12 weight tensors of the PPO Critic (ValueNetwork).
///
/// Layout matches `ValueNetwork` in `crates/ml/src/ppo/ppo.rs`:
/// - Layer 0: `[512, 54]`, `[512]`
/// - Layer 1: `[384, 512]`, `[384]`
/// - Layer 2: `[256, 384]`, `[256]`
/// - Layer 3: `[128, 256]`, `[128]`
/// - Layer 4: `[64, 128]`, `[64]`
/// - Output: `[1, 64]`, `[1]`
#[allow(missing_debug_implementations)]
pub struct PpoCriticWeightSet {
pub vw1: CudaSlice<f32>, // value_layer_0.weight [512, 54]
pub vb1: CudaSlice<f32>, // value_layer_0.bias [512]
pub vw2: CudaSlice<f32>, // value_layer_1.weight [384, 512]
pub vb2: CudaSlice<f32>, // value_layer_1.bias [384]
pub vw3: CudaSlice<f32>, // value_layer_2.weight [256, 384]
pub vb3: CudaSlice<f32>, // value_layer_2.bias [256]
pub vw4: CudaSlice<f32>, // value_layer_3.weight [128, 256]
pub vb4: CudaSlice<f32>, // value_layer_3.bias [128]
pub vw5: CudaSlice<f32>, // value_layer_4.weight [64, 128]
pub vb5: CudaSlice<f32>, // value_layer_4.bias [64]
pub vw6: CudaSlice<f32>, // value_output.weight [1, 64]
pub vb6: CudaSlice<f32>, // value_output.bias [1]
}
Step 3: Add extract/sync functions
Follow the extract_dueling_weights/sync_dueling_weights pattern:
/// Extract PPO actor weights from a `VarMap` and upload to GPU.
pub fn extract_ppo_actor_weights(
vars: &VarMap,
stream: &Arc<CudaStream>,
) -> Result<PpoActorWeightSet, MLError> {
let vars_data = vars.data().lock().map_err(|e| {
MLError::ModelError(format!("Failed to lock VarMap: {e}"))
})?;
let extract = |name: &str| extract_one(&vars_data, name, stream);
let pw1 = extract("policy_layer_0.weight")?;
let pb1 = extract("policy_layer_0.bias")?;
let pw2 = extract("policy_layer_1.weight")?;
let pb2 = extract("policy_layer_1.bias")?;
let pw3 = extract("policy_output.weight")?;
let pb3 = extract("policy_output.bias")?;
let total_params: usize = PPO_ACTOR_WEIGHT_NAMES.iter()
.filter_map(|name| vars_data.get(*name).map(|v| v.as_tensor().elem_count()))
.sum();
info!(total_params, "PPO actor weights extracted and uploaded to GPU");
Ok(PpoActorWeightSet { pw1, pb1, pw2, pb2, pw3, pb3 })
}
/// Extract PPO critic weights from a `VarMap` and upload to GPU.
pub fn extract_ppo_critic_weights(
vars: &VarMap,
stream: &Arc<CudaStream>,
) -> Result<PpoCriticWeightSet, MLError> {
let vars_data = vars.data().lock().map_err(|e| {
MLError::ModelError(format!("Failed to lock VarMap: {e}"))
})?;
let extract = |name: &str| extract_one(&vars_data, name, stream);
let vw1 = extract("value_layer_0.weight")?;
let vb1 = extract("value_layer_0.bias")?;
let vw2 = extract("value_layer_1.weight")?;
let vb2 = extract("value_layer_1.bias")?;
let vw3 = extract("value_layer_2.weight")?;
let vb3 = extract("value_layer_2.bias")?;
let vw4 = extract("value_layer_3.weight")?;
let vb4 = extract("value_layer_3.bias")?;
let vw5 = extract("value_layer_4.weight")?;
let vb5 = extract("value_layer_4.bias")?;
let vw6 = extract("value_output.weight")?;
let vb6 = extract("value_output.bias")?;
let total_params: usize = PPO_CRITIC_WEIGHT_NAMES.iter()
.filter_map(|name| vars_data.get(*name).map(|v| v.as_tensor().elem_count()))
.sum();
info!(total_params, "PPO critic weights extracted and uploaded to GPU");
Ok(PpoCriticWeightSet { vw1, vb1, vw2, vb2, vw3, vb3, vw4, vb4, vw5, vb5, vw6, vb6 })
}
/// Re-upload PPO actor weights into existing GPU buffers.
pub fn sync_ppo_actor_weights(
vars: &VarMap,
weights: &mut PpoActorWeightSet,
stream: &Arc<CudaStream>,
) -> Result<(), MLError> {
let vars_data = vars.data().lock().map_err(|e| {
MLError::ModelError(format!("Failed to lock VarMap: {e}"))
})?;
sync_one(&vars_data, "policy_layer_0.weight", &mut weights.pw1, stream)?;
sync_one(&vars_data, "policy_layer_0.bias", &mut weights.pb1, stream)?;
sync_one(&vars_data, "policy_layer_1.weight", &mut weights.pw2, stream)?;
sync_one(&vars_data, "policy_layer_1.bias", &mut weights.pb2, stream)?;
sync_one(&vars_data, "policy_output.weight", &mut weights.pw3, stream)?;
sync_one(&vars_data, "policy_output.bias", &mut weights.pb3, stream)?;
Ok(())
}
/// Re-upload PPO critic weights into existing GPU buffers.
pub fn sync_ppo_critic_weights(
vars: &VarMap,
weights: &mut PpoCriticWeightSet,
stream: &Arc<CudaStream>,
) -> Result<(), MLError> {
let vars_data = vars.data().lock().map_err(|e| {
MLError::ModelError(format!("Failed to lock VarMap: {e}"))
})?;
sync_one(&vars_data, "value_layer_0.weight", &mut weights.vw1, stream)?;
sync_one(&vars_data, "value_layer_0.bias", &mut weights.vb1, stream)?;
sync_one(&vars_data, "value_layer_1.weight", &mut weights.vw2, stream)?;
sync_one(&vars_data, "value_layer_1.bias", &mut weights.vb2, stream)?;
sync_one(&vars_data, "value_layer_2.weight", &mut weights.vw3, stream)?;
sync_one(&vars_data, "value_layer_2.bias", &mut weights.vb3, stream)?;
sync_one(&vars_data, "value_layer_3.weight", &mut weights.vw4, stream)?;
sync_one(&vars_data, "value_layer_3.bias", &mut weights.vb4, stream)?;
sync_one(&vars_data, "value_layer_4.weight", &mut weights.vw5, stream)?;
sync_one(&vars_data, "value_layer_4.bias", &mut weights.vb5, stream)?;
sync_one(&vars_data, "value_output.weight", &mut weights.vw6, stream)?;
sync_one(&vars_data, "value_output.bias", &mut weights.vb6, stream)?;
Ok(())
}
Step 4: Add unit tests
In the tests module of gpu_weights.rs, add:
use crate::ppo::ppo::{PPOConfig, PPO};
#[test]
fn test_ppo_actor_weight_key_paths() {
use super::PPO_ACTOR_WEIGHT_NAMES;
let config = PPOConfig {
state_dim: 54,
num_actions: 45,
policy_hidden_dims: vec![128, 64],
..PPOConfig::default()
};
let model = PPO::with_device(config, Device::Cpu).expect("PPO creation should not fail");
let vars_data = model.actor.vars().data().lock().expect("lock should succeed");
for &name in &PPO_ACTOR_WEIGHT_NAMES {
assert!(vars_data.contains_key(name), "Missing actor VarMap key: {name}");
}
let pw1 = vars_data.get("policy_layer_0.weight").unwrap().as_tensor();
assert_eq!(pw1.dims(), &[128, 54], "policy_layer_0.weight shape");
let pw2 = vars_data.get("policy_layer_1.weight").unwrap().as_tensor();
assert_eq!(pw2.dims(), &[64, 128], "policy_layer_1.weight shape");
let pw3 = vars_data.get("policy_output.weight").unwrap().as_tensor();
assert_eq!(pw3.dims(), &[45, 64], "policy_output.weight shape");
}
#[test]
fn test_ppo_critic_weight_key_paths() {
use super::PPO_CRITIC_WEIGHT_NAMES;
let config = PPOConfig {
state_dim: 54,
num_actions: 45,
value_hidden_dims: vec![512, 384, 256, 128, 64],
..PPOConfig::default()
};
let model = PPO::with_device(config, Device::Cpu).expect("PPO creation should not fail");
let vars_data = model.critic.vars().data().lock().expect("lock should succeed");
for &name in &PPO_CRITIC_WEIGHT_NAMES {
assert!(vars_data.contains_key(name), "Missing critic VarMap key: {name}");
}
let vw1 = vars_data.get("value_layer_0.weight").unwrap().as_tensor();
assert_eq!(vw1.dims(), &[512, 54], "value_layer_0.weight shape");
let vw5 = vars_data.get("value_layer_4.weight").unwrap().as_tensor();
assert_eq!(vw5.dims(), &[64, 128], "value_layer_4.weight shape");
let vw6 = vars_data.get("value_output.weight").unwrap().as_tensor();
assert_eq!(vw6.dims(), &[1, 64], "value_output.weight shape");
}
#[test]
fn test_ppo_actor_param_count() {
let config = PPOConfig {
state_dim: 54,
num_actions: 45,
policy_hidden_dims: vec![128, 64],
..PPOConfig::default()
};
let model = PPO::with_device(config, Device::Cpu).expect("PPO creation");
let vars_data = model.actor.vars().data().lock().expect("lock");
let total: usize = vars_data.values().map(|v| v.as_tensor().elem_count()).sum();
// 54*128+128 + 128*64+64 + 64*45+45 = 6912+128+8192+64+2880+45 = 18221
// Exact count depends on PPO implementation — assert non-zero and log
assert!(total > 15_000, "Actor params should be > 15K, got {total}");
assert!(total < 25_000, "Actor params should be < 25K, got {total}");
}
#[test]
fn test_ppo_critic_param_count() {
let config = PPOConfig {
state_dim: 54,
num_actions: 45,
value_hidden_dims: vec![512, 384, 256, 128, 64],
..PPOConfig::default()
};
let model = PPO::with_device(config, Device::Cpu).expect("PPO creation");
let vars_data = model.critic.vars().data().lock().expect("lock");
let total: usize = vars_data.values().map(|v| v.as_tensor().elem_count()).sum();
// 54*512+512 + 512*384+384 + 384*256+256 + 256*128+128 + 128*64+64 + 64*1+1
// = 27648+512+196608+384+98304+256+32768+128+8192+64+64+1 = ~364,929
assert!(total > 300_000, "Critic params should be > 300K, got {total}");
assert!(total < 400_000, "Critic params should be < 400K, got {total}");
}
Step 5: Run tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib gpu_weights -- --test-threads=1
Expected: 6/6 PASS (2 existing + 4 new)
Step 6: Commit
git add crates/ml/src/cuda_pipeline/gpu_weights.rs
git commit -m "feat(cuda): add PPO actor/critic weight extraction and GPU upload"
Task 6: PPO GPU experience collector (gpu_ppo_collector.rs)
Files:
- Create:
crates/ml/src/cuda_pipeline/gpu_ppo_collector.rs - Modify:
crates/ml/src/cuda_pipeline/mod.rs
Context: Mirror gpu_experience_collector.rs (DQN) but for PPO. Key differences: 6+12 actor/critic weights instead of 12+12 online/target, output buffers include log_probs, advantages, returns instead of target_q, td_error. Source compilation uses concatenation. GAE lambda is a scalar config arg.
Step 1: Create gpu_ppo_collector.rs
Follow the exact pattern from gpu_experience_collector.rs (lines 1-550+). The struct should be:
#![allow(unsafe_code)]
//! GPU-accelerated PPO experience collection via the zero-roundtrip kernel.
//!
//! Wraps `ppo_full_experience_kernel` from `ppo_experience_kernel.cu`, managing
//! all GPU buffers, compiling the kernel at runtime via NVRTC, and providing a
//! `collect_experiences()` method that launches the kernel and downloads results.
use std::sync::Arc;
use candle_core::cuda_backend::cudarc;
use cudarc::driver::{CudaFunction, CudaSlice, CudaStream, LaunchConfig, PushKernelArg};
use cudarc::nvrtc::Ptx;
use candle_nn::VarMap;
use tracing::{debug, info};
use crate::MLError;
use super::gpu_weights::{
CuriosityWeightSet, PpoActorWeightSet, PpoCriticWeightSet,
extract_ppo_actor_weights, extract_ppo_critic_weights, extract_curiosity_weights,
sync_ppo_actor_weights, sync_ppo_critic_weights, sync_curiosity_weights,
};
const STATE_DIM: usize = 54;
const MAX_EPISODES: usize = 256;
const MAX_TIMESTEPS: usize = 1000;
const PORTFOLIO_STATE_SIZE: usize = 8;
const BARRIER_STATE_SIZE: usize = 5;
const DIVERSITY_WINDOW: usize = 100;
Config struct:
#[derive(Debug, Clone)]
pub struct PpoCollectorConfig {
pub max_position: f32,
pub episode_length: i32,
pub total_bars: i32,
pub gamma: f32,
pub gae_lambda: f32,
pub curiosity_max_reward: f32,
pub n_episodes: i32,
pub timesteps_per_episode: i32,
pub barrier_profit_mult: f32,
pub barrier_loss_mult: f32,
pub barrier_max_bars: f32,
pub barrier_scale: f32,
pub diversity_scale: f32,
pub curiosity_scale: f32,
pub risk_weight: f32,
}
impl Default for PpoCollectorConfig {
fn default() -> Self {
Self {
max_position: 1.0,
episode_length: 500,
total_bars: 10_000,
gamma: 0.99,
gae_lambda: 0.95,
curiosity_max_reward: 0.1,
n_episodes: 128,
timesteps_per_episode: 500,
barrier_profit_mult: 1.02,
barrier_loss_mult: 0.98,
barrier_max_bars: 500.0,
barrier_scale: 0.5,
diversity_scale: 1.0,
curiosity_scale: 1.0,
risk_weight: 0.1,
}
}
}
impl PpoCollectorConfig {
pub fn total_experiences(&self) -> usize {
self.n_episodes as usize * self.timesteps_per_episode as usize
}
}
Output batch:
#[derive(Debug)]
pub struct PpoExperienceBatch {
pub states: Vec<f32>, // [N * L * STATE_DIM]
pub actions: Vec<i32>, // [N * L]
pub log_probs: Vec<f32>, // [N * L]
pub advantages: Vec<f32>, // [N * L]
pub returns: Vec<f32>, // [N * L]
pub done_flags: Vec<i32>, // [N * L]
pub n_episodes: usize,
pub timesteps: usize,
}
Main struct following GpuExperienceCollector pattern:
#[allow(missing_debug_implementations)]
pub struct GpuPpoExperienceCollector {
stream: Arc<CudaStream>,
kernel_func: CudaFunction,
// PPO network weights
actor_weights: PpoActorWeightSet,
critic_weights: PpoCriticWeightSet,
curiosity_weights: CuriosityWeightSet,
// Per-episode state buffers
portfolio_states: CudaSlice<f32>,
barrier_states: CudaSlice<f32>,
diversity_windows: CudaSlice<i32>,
diversity_metas: CudaSlice<i32>,
barrier_config: CudaSlice<f32>,
rng_states: CudaSlice<u32>,
episode_starts_buf: CudaSlice<i32>,
// Output buffers
states_out: CudaSlice<f32>,
actions_out: CudaSlice<i32>,
log_probs_out: CudaSlice<f32>,
advantages_out: CudaSlice<f32>,
returns_out: CudaSlice<f32>,
dones_out: CudaSlice<i32>,
}
Constructor new(): compile kernel via source concatenation, extract weights, allocate buffers. Key difference from DQN:
let common_src = include_str!("common_device_functions.cuh");
let kernel_src = include_str!("ppo_experience_kernel.cu");
let full_source = format!("{}\n{}", common_src, kernel_src);
let ptx: Ptx = cudarc::nvrtc::compile_ptx(&full_source).map_err(|e| {
MLError::ModelError(format!("CUDA ppo_experience_kernel compilation failed: {e}"))
})?;
// ... load_module → load_function("ppo_full_experience_kernel")
collect_experiences(): launch kernel with args matching the kernel signature (3 data pointers + 6 actor + 12 critic + 4 curiosity + 4 state + 1 barrier_config + 12 scalars + 1 rng + 6 outputs = 49 args), then download 6 output buffers.
sync_weights(): call sync_ppo_actor_weights, sync_ppo_critic_weights, sync_curiosity_weights.
Step 2: Add module declaration
In crates/ml/src/cuda_pipeline/mod.rs, add:
#[cfg(feature = "cuda")]
pub mod gpu_ppo_collector;
Step 3: Run compile check
Run: SQLX_OFFLINE=true cargo check -p ml --lib
Expected: Clean compile (0 errors, 0 warnings)
Step 4: Commit
git add crates/ml/src/cuda_pipeline/gpu_ppo_collector.rs \
crates/ml/src/cuda_pipeline/mod.rs
git commit -m "feat(cuda): add GpuPpoExperienceCollector with zero-roundtrip kernel launch"
Task 7: Unit tests for PPO weight extraction and kernel source
Files:
- Modify:
crates/ml/src/cuda_pipeline/mod.rs
Context: Add CPU-only unit tests that verify PPO kernel source, config defaults, and weight key paths.
Step 1: Add tests to mod.rs
#[cfg(feature = "cuda")]
#[test]
fn test_ppo_collector_config_defaults() {
use super::gpu_ppo_collector::PpoCollectorConfig;
let cfg = PpoCollectorConfig::default();
assert_eq!(cfg.n_episodes, 128);
assert_eq!(cfg.timesteps_per_episode, 500);
assert_eq!(cfg.total_experiences(), 64_000);
assert!((cfg.gamma - 0.99).abs() < f32::EPSILON);
assert!((cfg.gae_lambda - 0.95).abs() < f32::EPSILON);
}
#[test]
fn test_ppo_kernel_source_concatenation() {
let common = include_str!("common_device_functions.cuh");
let kernel = include_str!("ppo_experience_kernel.cu");
let full = format!("{}\n{}", common, kernel);
// The concatenated source should have both shared functions and PPO kernel
assert!(full.contains("gpu_random"));
assert!(full.contains("matvec_leaky_relu"));
assert!(full.contains("ppo_full_experience_kernel"));
assert!(full.contains("compute_gae_backward"));
assert!(full.contains("softmax_sample"));
// Should NOT contain DQN-specific functions
assert!(!full.contains("q_forward_dueling"));
}
Step 2: Run all tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib cuda_pipeline -- --test-threads=1
Expected: All PASS (existing + new)
Run: SQLX_OFFLINE=true cargo test -p ml --lib gpu_weights -- --test-threads=1
Expected: All PASS
Step 3: Commit
git add crates/ml/src/cuda_pipeline/mod.rs
git commit -m "test(cuda): add PPO kernel source verification and config default tests"
Task 8: PPO trainer integration
Files:
- Modify:
crates/ml/src/trainers/ppo.rs
Context: Add gpu_ppo_collector: Option<GpuPpoExperienceCollector> field to PpoTrainer behind #[cfg(feature = "cuda")]. Initialize when CUDA device present and curiosity available. The existing collect_rollouts() CPU path remains as fallback. Weight sync after model.update().
The PPO model is held behind Arc<Mutex<PPO>> — we need to access model.actor.vars() and model.critic.vars() for weight extraction. The PPO struct has public actor: PolicyNetwork and critic: ValueNetwork fields with .vars() methods.
Step 1: Add field to PpoTrainer
After explained_variance_history:
#[cfg(feature = "cuda")]
gpu_ppo_collector: Option<crate::cuda_pipeline::gpu_ppo_collector::GpuPpoExperienceCollector>,
Step 2: Initialize in constructor
In PpoTrainer::new(), after the Ok(Self { ... }) block, add initialization:
#[cfg(feature = "cuda")]
gpu_ppo_collector: None,
Then in train(), before the main loop, add a GPU initialization block:
// Phase 2c: Initialize GPU PPO experience collector if CUDA available
#[cfg(feature = "cuda")]
let mut gpu_ppo_collector: Option<crate::cuda_pipeline::gpu_ppo_collector::GpuPpoExperienceCollector> = None;
#[cfg(feature = "cuda")]
{
if self.device.is_cuda() {
use candle_core::cuda_backend::cudarc;
let model = self.model.lock().await;
let actor_vars = model.actor.vars();
let critic_vars = model.critic.vars();
// For curiosity, check if available (PPO doesn't have built-in curiosity yet)
// Use a dummy VarMap for now — curiosity integration is optional
match (|| -> Result<_, MLError> {
let cuda_device = match &self.device {
Device::Cuda(d) => d,
_ => return Err(MLError::ModelError("Not a CUDA device".into())),
};
let stream = std::sync::Arc::new(cuda_device.cuda_stream());
crate::cuda_pipeline::gpu_ppo_collector::GpuPpoExperienceCollector::new(
stream,
actor_vars,
critic_vars,
// Curiosity vars — PPO doesn't have curiosity module yet,
// so this will need a separate VarMap or be made optional
&candle_nn::VarMap::new(), // placeholder
1_000_000.0, // initial_capital
0.0001, // avg_spread
0.20, // cash_reserve_pct
)
})() {
Ok(collector) => {
info!("PPO GPU experience collector initialized");
gpu_ppo_collector = Some(collector);
}
Err(e) => {
warn!("PPO GPU collector init failed, using CPU path: {e}");
}
}
}
}
Step 3: Add weight sync after update
After model.update(&mut training_batch)? (around line 352), add:
// Phase 2c: Sync weights to GPU after PPO update
#[cfg(feature = "cuda")]
if let Some(ref mut collector) = gpu_ppo_collector {
let model = self.model.lock().await;
if let Err(e) = collector.sync_weights(model.actor.vars(), model.critic.vars()) {
warn!("GPU weight sync failed: {e}");
}
}
Note: The sync_weights method on GpuPpoExperienceCollector needs to accept actor and critic VarMaps separately and call sync_ppo_actor_weights + sync_ppo_critic_weights + optionally sync_curiosity_weights.
Step 4: Run compile check
Run: SQLX_OFFLINE=true cargo check -p ml --lib
Expected: Clean compile
Step 5: Run PPO tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib ppo -- --test-threads=1
Expected: All existing PPO tests pass (CPU fallback)
Step 6: Commit
git add crates/ml/src/trainers/ppo.rs \
crates/ml/src/cuda_pipeline/gpu_ppo_collector.rs
git commit -m "feat(ppo): integrate GPU experience collector with weight sync"
Task 9: DQN regression verification
Files: None (read-only verification)
Context: Verify the shared header extraction didn't break any DQN functionality.
Step 1: Run DQN tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib dqn -- --test-threads=1
Expected: 427/427 PASS (or current count)
Step 2: Run cuda_pipeline tests
Run: SQLX_OFFLINE=true cargo test -p ml --lib cuda_pipeline -- --test-threads=1
Expected: All PASS
Step 3: Run full ML suite
Run: SQLX_OFFLINE=true cargo test -p ml --lib -- --test-threads=1
Expected: 2396+ PASS
Task 10: Workspace verification
Files: None (read-only verification)
Step 1: Workspace compile check
Run: SQLX_OFFLINE=true cargo check --workspace
Expected: Clean compile
Step 2: Clippy
Run: SQLX_OFFLINE=true cargo clippy -p ml --lib -- -D warnings
Expected: 0 warnings
Step 3: Full ML test suite
Run: SQLX_OFFLINE=true cargo test -p ml --lib -- --test-threads=1
Expected: 2396+ PASS (all existing + new tests)
Step 4: Commit if any outstanding changes
Task 11: Final summary
Review commit history for Phase 2c, verify all files are committed, count lines added/modified.
Expected commits:
refactor(cuda): extract shared device functions to common_device_functions.cuhfeat(cuda): add PPO actor forward and softmax sampling device functionsfeat(cuda): add PPO critic forward (5-layer ping-pong) and GAE backward scanfeat(cuda): add PPO main kernel — forward rollout + backward GAE scanfeat(cuda): add PPO actor/critic weight extraction and GPU uploadfeat(cuda): add GpuPpoExperienceCollector with zero-roundtrip kernel launchtest(cuda): add PPO kernel source verification and config default testsfeat(ppo): integrate GPU experience collector with weight sync
Files created (3):
crates/ml/src/cuda_pipeline/common_device_functions.cuhcrates/ml/src/cuda_pipeline/ppo_experience_kernel.cucrates/ml/src/cuda_pipeline/gpu_ppo_collector.rs
Files modified (4):
crates/ml/src/cuda_pipeline/dqn_experience_kernel.cu(shared functions extracted)crates/ml/src/cuda_pipeline/gpu_experience_collector.rs(source concatenation)crates/ml/src/cuda_pipeline/gpu_weights.rs(PPO weight sets)crates/ml/src/cuda_pipeline/mod.rs(module declarations + tests)crates/ml/src/trainers/ppo.rs(GPU collector integration)