Files
foxhunt/docs/plans/2026-02-28-gpu-pipeline-phase2c-implementation.md
2026-02-28 13:07:20 +01:00

55 KiB
Raw Blame History

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:

  1. refactor(cuda): extract shared device functions to common_device_functions.cuh
  2. feat(cuda): add PPO actor forward and softmax sampling device functions
  3. feat(cuda): add PPO critic forward (5-layer ping-pong) and GAE backward scan
  4. feat(cuda): add PPO main kernel — forward rollout + backward GAE scan
  5. feat(cuda): add PPO actor/critic weight extraction and GPU upload
  6. feat(cuda): add GpuPpoExperienceCollector with zero-roundtrip kernel launch
  7. test(cuda): add PPO kernel source verification and config default tests
  8. feat(ppo): integrate GPU experience collector with weight sync

Files created (3):

  • crates/ml/src/cuda_pipeline/common_device_functions.cuh
  • crates/ml/src/cuda_pipeline/ppo_experience_kernel.cu
  • crates/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)