From 612e1190f6c19570f8d2f3381f6787f1a7d99dab Mon Sep 17 00:00:00 2001 From: haixuantao Date: Tue, 26 May 2026 16:20:42 +0100 Subject: [PATCH 1/2] Add tanh activation + Adam optimizer GPU ops Adds vortx::linalg::{Activation (tanh + tanh_backward), Adam} and their shaders, the GPU building blocks for MLP training (used by nexus RL demos / zealot-rl). Co-Authored-By: Claude Opus 4.7 (1M context) --- src/linalg/activation.rs | 72 ++++++++++++++++++++++++++ src/linalg/mod.rs | 4 ++ src/linalg/optim.rs | 61 ++++++++++++++++++++++ vortx-shaders/Cargo.toml | 4 ++ vortx-shaders/src/linalg/activation.rs | 54 +++++++++++++++++++ vortx-shaders/src/linalg/mod.rs | 6 +++ vortx-shaders/src/linalg/optim.rs | 63 ++++++++++++++++++++++ 7 files changed, 264 insertions(+) create mode 100644 src/linalg/activation.rs create mode 100644 src/linalg/optim.rs create mode 100644 vortx-shaders/src/linalg/activation.rs create mode 100644 vortx-shaders/src/linalg/optim.rs diff --git a/src/linalg/activation.rs b/src/linalg/activation.rs new file mode 100644 index 0000000..91867db --- /dev/null +++ b/src/linalg/activation.rs @@ -0,0 +1,72 @@ +//! Element-wise activation functions (host dispatch). +//! +//! Added for zealot's MLP policy — vortx upstream has no activations. + +use crate::shaders::linalg::{GpuTanh, GpuTanhBackward}; +use crate::shapes::TensorLayoutBuffers; +use crate::tensor::{AsTensorMut, AsTensorRef}; +use khal::Shader; +use khal::backend::{GpuBackend, GpuBackendError, GpuPass}; + +/// Element-wise activation kernels. +#[derive(Shader)] +pub struct Activation { + /// In-place tanh. + pub tanh: GpuTanh, + /// In-place tanh backward (`g *= 1 - y^2`). + pub tanh_backward: GpuTanhBackward, +} + +impl Activation { + /// In-place tanh: `a = tanh(a)`. + pub fn tanh( + &self, + backend: &GpuBackend, + shapes: &mut TensorLayoutBuffers, + pass: &mut GpuPass, + mut a: impl AsTensorMut, + ) -> Result<(), GpuBackendError> { + let mut a = a.as_tensor_mut(); + let shape_a = a.layout().canonicalize(); + let num_threads = a.len() as u32; + + shapes.insert(backend, shape_a)?; + let shape_a_buf = shapes.get(shape_a).unwrap(); + let mut buf_a = a.buffer_mut(); + + self.tanh + .call(pass, num_threads, &shape_a_buf.as_slice(), &mut buf_a) + } + + /// In-place tanh backward: `g *= 1 - y^2`, where `y = tanh(x)` is the forward output. + /// `g` and `y` must have the same shape. + pub fn tanh_backward( + &self, + backend: &GpuBackend, + shapes: &mut TensorLayoutBuffers, + pass: &mut GpuPass, + mut g: impl AsTensorMut, + y: impl AsTensorRef, + ) -> Result<(), GpuBackendError> { + let mut g = g.as_tensor_mut(); + let y = y.as_tensor_ref(); + let shape_g = g.layout().canonicalize(); + let shape_y = y.layout().canonicalize(); + let num_threads = g.len() as u32; + + shapes.insert(backend, shape_g)?; + shapes.insert(backend, shape_y)?; + let shape_g_buf = shapes.get(shape_g).unwrap(); + let shape_y_buf = shapes.get(shape_y).unwrap(); + let mut buf_g = g.buffer_mut(); + + self.tanh_backward.call( + pass, + num_threads, + &shape_g_buf.as_slice(), + &shape_y_buf.as_slice(), + &mut buf_g, + &y.buffer(), + ) + } +} diff --git a/src/linalg/mod.rs b/src/linalg/mod.rs index 7a65987..c13ed33 100644 --- a/src/linalg/mod.rs +++ b/src/linalg/mod.rs @@ -1,14 +1,18 @@ //! Fundamental linear-algebra matrix/vector operations. +mod activation; mod contiguous; mod gemm; mod op_assign; +mod optim; mod reduce; mod repeat; +pub use activation::Activation; pub use contiguous::Contiguous; pub use gemm::{Gemm, MatrixMode, N, T}; pub use op_assign::{BinOpOffsets, OpAssign, OpAssignVariant}; +pub use optim::{Adam, AdamParams}; pub use reduce::{Reduce, ReduceVariant}; pub use repeat::Repeat; diff --git a/src/linalg/optim.rs b/src/linalg/optim.rs new file mode 100644 index 0000000..3c9b98e --- /dev/null +++ b/src/linalg/optim.rs @@ -0,0 +1,61 @@ +//! Optimizer host dispatch (Adam). Added for zealot. + +use crate::shaders::linalg::GpuAdam; +use crate::shapes::TensorLayoutBuffers; +use crate::tensor::{AsTensorMut, AsTensorRef}; +use khal::Shader; +use khal::backend::{GpuBackend, GpuBackendError, GpuPass}; + +// Re-export the params struct from the shader crate. +pub use vortx_shaders::linalg::optim::AdamParams; + +/// The Adam optimizer kernel. +#[derive(Shader)] +pub struct Adam { + /// One in-place Adam update step. + pub adam: GpuAdam, +} + +impl Adam { + /// Performs one in-place Adam step: updates `theta`, `m`, `v` from `grad`. + /// + /// `params` is a scalar `Tensor` (UNIFORM usage); `theta`, `grad`, + /// `m`, `v` all share the same shape. + pub fn step( + &self, + backend: &GpuBackend, + shapes: &mut TensorLayoutBuffers, + pass: &mut GpuPass, + params: impl AsTensorRef, + mut theta: impl AsTensorMut, + grad: impl AsTensorRef, + mut m: impl AsTensorMut, + mut v: impl AsTensorMut, + ) -> Result<(), GpuBackendError> { + let params = params.as_tensor_ref(); + let mut theta = theta.as_tensor_mut(); + let grad = grad.as_tensor_ref(); + let mut m = m.as_tensor_mut(); + let mut v = v.as_tensor_mut(); + + let shape = theta.layout().canonicalize(); + let num_threads = theta.len() as u32; + + shapes.insert(backend, shape)?; + let shape_buf = shapes.get(shape).unwrap(); + let mut buf_theta = theta.buffer_mut(); + let mut buf_m = m.buffer_mut(); + let mut buf_v = v.buffer_mut(); + + self.adam.call( + pass, + num_threads, + &shape_buf.as_slice(), + ¶ms.buffer(), + &mut buf_theta, + &grad.buffer(), + &mut buf_m, + &mut buf_v, + ) + } +} diff --git a/vortx-shaders/Cargo.toml b/vortx-shaders/Cargo.toml index 998d75f..ebbf2f1 100644 --- a/vortx-shaders/Cargo.toml +++ b/vortx-shaders/Cargo.toml @@ -27,6 +27,10 @@ cuda = ["khal-std/cuda", "khal/cuda"] khal-std = { workspace = true } # glamx provides UVec3 and other glam types (no_std compatible, used on all targets). glamx = { version = "0.2", default-features = false, features = ["nostd-libm", "bytemuck"] } +# Cap glam < 0.33 for the shader build: spirv-std 0.10.0-alpha.1 declares glam ">=0.30.8" +# (open-ended), but glam 0.33 dropped `UVec4` under default-features=false, which breaks +# spirv-std's compile (112 errors). 0.32.1 is the newest pre-0.33 version that still works. +glam = { version = "=0.32.1", default-features = false } # Host-only dependencies (excluded on GPU targets: spirv and nvptx64). [target.'cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))'.dependencies] diff --git a/vortx-shaders/src/linalg/activation.rs b/vortx-shaders/src/linalg/activation.rs new file mode 100644 index 0000000..04c880d --- /dev/null +++ b/vortx-shaders/src/linalg/activation.rs @@ -0,0 +1,54 @@ +//! Element-wise activation functions (tanh forward/backward). +//! +//! vortx upstream has no activations; these were added for zealot's MLP policy. +//! Uniform-shape bindings only (no push_constants variant), matching the default build. + +use super::shape::Shape; +use crate::utils::limits::MAX_NUM_WORKGROUPS; +use crate::utils::trig::stable_tanh; +use glamx::UVec3; +use khal_std::{ + index::MaybeIndexUnchecked, + macros::{spirv, spirv_bindgen}, +}; + +const WORKGROUP_SIZE: u32 = 256; +const MAX_NUM_THREADS: u32 = MAX_NUM_WORKGROUPS * WORKGROUP_SIZE; + +/// Element-wise tanh, in place: `a = tanh(a)`. +#[spirv_bindgen] +#[spirv(compute(threads(256, 1, 1)))] +pub fn gpu_tanh( + #[spirv(global_invocation_id)] invocation_id: UVec3, + #[spirv(uniform, descriptor_set = 0, binding = 0)] shape_a: &Shape, + #[spirv(storage_buffer, descriptor_set = 0, binding = 1)] a: &mut [f32], +) { + for thread_id in (invocation_id.x..shape_a.len()).step_by(MAX_NUM_THREADS as usize) { + let id = shape_a.decompose(thread_id); + let ia = shape_a.it_vec(id) as usize; + let slot = a.at_mut(ia); + *slot = stable_tanh(*slot); + } +} + +/// Backward of tanh, in place: `g *= 1 - y*y`, where `y = tanh(x)` is the forward output. +/// +/// `g` and `y` are expected to have the same shape (the per-element local derivative +/// of tanh is `1 - tanh(x)^2`, expressed in terms of the cached output `y`). +#[spirv_bindgen] +#[spirv(compute(threads(256, 1, 1)))] +pub fn gpu_tanh_backward( + #[spirv(global_invocation_id)] invocation_id: UVec3, + #[spirv(uniform, descriptor_set = 0, binding = 0)] shape_g: &Shape, + #[spirv(uniform, descriptor_set = 0, binding = 1)] shape_y: &Shape, + #[spirv(storage_buffer, descriptor_set = 0, binding = 2)] g: &mut [f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 3)] y: &[f32], +) { + for thread_id in (invocation_id.x..shape_g.len()).step_by(MAX_NUM_THREADS as usize) { + let id = shape_g.decompose(thread_id); + let ig = shape_g.it_vec(id) as usize; + let iy = shape_y.it_vec(id) as usize; + let yi = y.read(iy); + *g.at_mut(ig) *= 1.0 - yi * yi; + } +} diff --git a/vortx-shaders/src/linalg/mod.rs b/vortx-shaders/src/linalg/mod.rs index dd0b36b..eeb1305 100644 --- a/vortx-shaders/src/linalg/mod.rs +++ b/vortx-shaders/src/linalg/mod.rs @@ -1,9 +1,11 @@ //! Linear algebra modules for shaders. +pub mod activation; pub mod contiguous; pub mod gemm; pub mod inv; pub mod op_assign; +pub mod optim; pub mod reduce; pub mod repeat; pub mod shape; @@ -14,12 +16,16 @@ pub use shape::{Shapes1, Shapes2, Shapes3}; // Re-export generated ShaderArgs structs (only available on host) #[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] +pub use activation::{GpuTanh, GpuTanhBackward}; +#[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] pub use contiguous::{Contiguous, ContiguousWithOffset}; #[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] pub use gemm::{GemmNaive, GemmTiled}; #[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] pub use op_assign::{GpuAdd, GpuCopy, GpuCopyWithOffsets, GpuDiv, GpuMul, GpuSub}; #[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] +pub use optim::GpuAdam; +#[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] pub use reduce::{ReduceAdd, ReduceMax, ReduceMin, ReduceMul, ReduceSqNorm}; #[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] pub use repeat::Repeat; diff --git a/vortx-shaders/src/linalg/optim.rs b/vortx-shaders/src/linalg/optim.rs new file mode 100644 index 0000000..da6b6df --- /dev/null +++ b/vortx-shaders/src/linalg/optim.rs @@ -0,0 +1,63 @@ +//! Optimizer kernels (Adam). Added for zealot; vortx upstream has no optimizers. + +use super::shape::Shape; +use crate::utils::limits::MAX_NUM_WORKGROUPS; +use glamx::UVec3; +use khal_std::{ + index::MaybeIndexUnchecked, + macros::{spirv, spirv_bindgen}, +}; +#[cfg(any(target_arch = "spirv", target_arch = "nvptx64"))] +use khal_std::num_traits::Float; + +const WORKGROUP_SIZE: u32 = 256; +const MAX_NUM_THREADS: u32 = MAX_NUM_WORKGROUPS * WORKGROUP_SIZE; + +/// Scalar parameters for one Adam step (uniform buffer; padded to 32 bytes). +#[repr(C)] +#[derive(Clone, Copy)] +#[cfg_attr( + not(any(target_arch = "spirv", target_arch = "nvptx64")), + derive(bytemuck::Pod, bytemuck::Zeroable) +)] +pub struct AdamParams { + pub lr: f32, + pub beta1: f32, + pub beta2: f32, + pub eps: f32, + /// `1 - beta1^t` (bias correction for the first moment). + pub bias_correction1: f32, + /// `1 - beta2^t` (bias correction for the second moment). + pub bias_correction2: f32, + pub pad0: f32, + pub pad1: f32, +} + +/// One in-place Adam step: updates first/second moments `m`, `v` and parameters +/// `theta` from the gradient `grad`. All buffers share `theta`'s shape. +#[spirv_bindgen] +#[spirv(compute(threads(256, 1, 1)))] +pub fn gpu_adam( + #[spirv(global_invocation_id)] invocation_id: UVec3, + #[spirv(uniform, descriptor_set = 0, binding = 0)] shape: &Shape, + #[spirv(uniform, descriptor_set = 0, binding = 1)] params: &AdamParams, + #[spirv(storage_buffer, descriptor_set = 0, binding = 2)] theta: &mut [f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 3)] grad: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 4)] m: &mut [f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 5)] v: &mut [f32], +) { + for thread_id in (invocation_id.x..shape.len()).step_by(MAX_NUM_THREADS as usize) { + let id = shape.decompose(thread_id); + let i = shape.it_vec(id) as usize; + let g = grad.read(i); + let m_old = *m.at_mut(i); + let v_old = *v.at_mut(i); + let mi = params.beta1 * m_old + (1.0 - params.beta1) * g; + let vi = params.beta2 * v_old + (1.0 - params.beta2) * g * g; + *m.at_mut(i) = mi; + *v.at_mut(i) = vi; + let mhat = mi / params.bias_correction1; + let vhat = vi / params.bias_correction2; + *theta.at_mut(i) -= params.lr * mhat / (vhat.sqrt() + params.eps); + } +} From 9214d98fd78e2dd7bee27b2eec732c68e6e5ac48 Mon Sep 17 00:00:00 2001 From: haixuantao Date: Fri, 5 Jun 2026 10:22:21 +0200 Subject: [PATCH 2/2] PPO loss-gradient GPU kernels New `ppo` op (host + shader) producing the per-sample OUTPUT gradients that feed the generic GEMM/elu_backward backward backbone: - gpu_ppo_actor_grad: clipped-surrogate actor gradient (logp over the action dims, ratio = exp(logp - logp_old), clip mask) -> g_mean plus the state-independent log_std gradient contribution. - gpu_ppo_value_grad: clipped value-loss gradient. Both are an exact port of zealot-rl's minibatch_step. Every per-sample tensor is row-major [rows x M] (M = minibatch columns); one thread handles one sample column and loops over the (small) action dim. No Shape uniform -- dims ride in PpoActorParams/PpoValueParams. Exports Ppo, PpoActorParams, PpoValueParams (host) and GpuPpoActorGrad, GpuPpoValueGrad (shaders). Verified vs CPU minibatch_step (~1e-7, ~25% of samples on the clip branch). Note: the one-line glamx Cargo.toml dependency is shared with the ELU and GEMM-vec4 branches; whichever lands first, the others need a trivial rebase of that line. Co-Authored-By: Claude Opus 4.8 (1M context) --- Cargo.toml | 1 + src/linalg/mod.rs | 2 + src/linalg/ppo.rs | 99 +++++++++++++++++++ vortx-shaders/src/linalg/mod.rs | 3 + vortx-shaders/src/linalg/ppo.rs | 167 ++++++++++++++++++++++++++++++++ 5 files changed, 272 insertions(+) create mode 100644 src/linalg/ppo.rs create mode 100644 vortx-shaders/src/linalg/ppo.rs diff --git a/Cargo.toml b/Cargo.toml index c888def..c7b3aac 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -32,6 +32,7 @@ khal = { version = "0.1", features = ["derive"]} [dependencies] bytemuck = "1" +glamx = { version = "0.2", default-features = false, features = ["bytemuck"] } include_dir = "0.7" nalgebra = "0.34" khal = { workspace = true } diff --git a/src/linalg/mod.rs b/src/linalg/mod.rs index c13ed33..7afe069 100644 --- a/src/linalg/mod.rs +++ b/src/linalg/mod.rs @@ -5,6 +5,7 @@ mod contiguous; mod gemm; mod op_assign; mod optim; +mod ppo; mod reduce; mod repeat; @@ -13,6 +14,7 @@ pub use contiguous::Contiguous; pub use gemm::{Gemm, MatrixMode, N, T}; pub use op_assign::{BinOpOffsets, OpAssign, OpAssignVariant}; pub use optim::{Adam, AdamParams}; +pub use ppo::{Ppo, PpoActorParams, PpoValueParams}; pub use reduce::{Reduce, ReduceVariant}; pub use repeat::Repeat; diff --git a/src/linalg/ppo.rs b/src/linalg/ppo.rs new file mode 100644 index 0000000..069555f --- /dev/null +++ b/src/linalg/ppo.rs @@ -0,0 +1,99 @@ +//! PPO loss-gradient host dispatch. Added for zealot's GPU policy update. +//! +//! Wraps the two PPO output-gradient kernels (clipped-surrogate actor gradient + +//! log_std contribution, and clipped value-loss gradient). These have no `Shape` +//! uniform — dimensions ride in the params struct and indexing is row-major — so +//! no `TensorLayoutBuffers` is needed. + +use crate::shaders::linalg::{GpuPpoActorGrad, GpuPpoValueGrad}; +use crate::tensor::{AsTensorMut, AsTensorRef}; +use khal::Shader; +use khal::backend::{GpuBackend, GpuBackendError, GpuPass}; + +// Re-export the params structs from the shader crate. +pub use vortx_shaders::linalg::ppo::{PpoActorParams, PpoValueParams}; + +/// PPO loss-gradient kernels. +#[derive(Shader)] +pub struct Ppo { + /// Clipped-surrogate actor gradient + log_std contribution. + pub actor_grad: GpuPpoActorGrad, + /// Clipped value-loss gradient. + pub value_grad: GpuPpoValueGrad, +} + +impl Ppo { + /// Actor PPO gradient. All per-sample tensors are row-major `[action_dim x M]` + /// except `log_std` (`[action_dim]`), `adv` / `logp_old` (`[M]`). Writes + /// `g_mean` and `g_logstd` (`[action_dim x M]`). `params.num_cols` must equal `M`. + #[allow(clippy::too_many_arguments)] + pub fn actor_grad( + &self, + pass: &mut GpuPass, + params: impl AsTensorRef, + mean: impl AsTensorRef, + action: impl AsTensorRef, + log_std: impl AsTensorRef, + adv: impl AsTensorRef, + logp_old: impl AsTensorRef, + mut g_mean: impl AsTensorMut, + mut g_logstd: impl AsTensorMut, + ) -> Result<(), GpuBackendError> { + let params = params.as_tensor_ref(); + let mean = mean.as_tensor_ref(); + let action = action.as_tensor_ref(); + let log_std = log_std.as_tensor_ref(); + let adv = adv.as_tensor_ref(); + let logp_old = logp_old.as_tensor_ref(); + let mut g_mean = g_mean.as_tensor_mut(); + let mut g_logstd = g_logstd.as_tensor_mut(); + + let num_threads = adv.len() as u32; // one thread per sample column + let mut buf_g_mean = g_mean.buffer_mut(); + let mut buf_g_logstd = g_logstd.buffer_mut(); + + self.actor_grad.call( + pass, + num_threads, + ¶ms.buffer(), + &mean.buffer(), + &action.buffer(), + &log_std.buffer(), + &adv.buffer(), + &logp_old.buffer(), + &mut buf_g_mean, + &mut buf_g_logstd, + ) + } + + /// Clipped value-loss gradient. `v_pred` / `value_old` / `ret` are `[M]`; + /// writes `g_v` (`[M]`). `params.num_cols` must equal `M`. + pub fn value_grad( + &self, + pass: &mut GpuPass, + params: impl AsTensorRef, + v_pred: impl AsTensorRef, + value_old: impl AsTensorRef, + ret: impl AsTensorRef, + mut g_v: impl AsTensorMut, + ) -> Result<(), GpuBackendError> { + let params = params.as_tensor_ref(); + let v_pred = v_pred.as_tensor_ref(); + let value_old = value_old.as_tensor_ref(); + let ret = ret.as_tensor_ref(); + let mut g_v = g_v.as_tensor_mut(); + + let num_threads = v_pred.len() as u32; + let mut buf_g_v = g_v.buffer_mut(); + + self.value_grad.call( + pass, + num_threads, + ¶ms.buffer(), + &v_pred.buffer(), + &value_old.buffer(), + &ret.buffer(), + &mut buf_g_v, + ) + } +} diff --git a/vortx-shaders/src/linalg/mod.rs b/vortx-shaders/src/linalg/mod.rs index eeb1305..8bda8b0 100644 --- a/vortx-shaders/src/linalg/mod.rs +++ b/vortx-shaders/src/linalg/mod.rs @@ -6,6 +6,7 @@ pub mod gemm; pub mod inv; pub mod op_assign; pub mod optim; +pub mod ppo; pub mod reduce; pub mod repeat; pub mod shape; @@ -26,6 +27,8 @@ pub use op_assign::{GpuAdd, GpuCopy, GpuCopyWithOffsets, GpuDiv, GpuMul, GpuSub} #[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] pub use optim::GpuAdam; #[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] +pub use ppo::{GpuPpoActorGrad, GpuPpoValueGrad}; +#[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] pub use reduce::{ReduceAdd, ReduceMax, ReduceMin, ReduceMul, ReduceSqNorm}; #[cfg(not(any(target_arch = "spirv", target_arch = "nvptx64")))] pub use repeat::Repeat; diff --git a/vortx-shaders/src/linalg/ppo.rs b/vortx-shaders/src/linalg/ppo.rs new file mode 100644 index 0000000..f538f82 --- /dev/null +++ b/vortx-shaders/src/linalg/ppo.rs @@ -0,0 +1,167 @@ +//! PPO loss-gradient kernels (added for zealot's GPU policy update). +//! +//! These produce the per-sample OUTPUT gradients that feed the generic +//! GEMM/`elu_backward` backward backbone: the clipped-surrogate actor gradient +//! `g_mean` plus the state-independent `log_std` gradient contribution, and the +//! clipped value-loss gradient. An exact port of `zealot-rl`'s `minibatch_step` +//! (ppo.rs). Every per-sample tensor is row-major `[rows x M]` (M = minibatch +//! columns); one GPU thread handles one sample column `m`, looping over the +//! (small) action dimension internally. + +use crate::utils::limits::MAX_NUM_WORKGROUPS; +use glamx::UVec3; +use khal_std::{ + index::MaybeIndexUnchecked, + macros::{spirv, spirv_bindgen}, +}; +#[cfg(any(target_arch = "spirv", target_arch = "nvptx64"))] +use khal_std::num_traits::Float; + +const WORKGROUP_SIZE: u32 = 256; +const MAX_NUM_THREADS: u32 = MAX_NUM_WORKGROUPS * WORKGROUP_SIZE; + +/// Scalar parameters for the actor PPO gradient (uniform buffer; 32 bytes). +#[repr(C)] +#[derive(Clone, Copy)] +#[cfg_attr( + not(any(target_arch = "spirv", target_arch = "nvptx64")), + derive(bytemuck::Pod, bytemuck::Zeroable) +)] +pub struct PpoActorParams { + /// PPO clip epsilon. + pub clip: f32, + /// Entropy bonus coefficient (subtracted from the log_std gradient). + pub entropy_coef: f32, + /// Per-sample averaging factor `1 / minibatch_size`. + pub scale: f32, + /// `0.5·ln(2π)` — the Gaussian log-prob normalisation constant. + pub log_sqrt_2pi: f32, + /// Action dimensionality (rows). + pub action_dim: u32, + /// Number of sample columns `M`. + pub num_cols: u32, + pub pad0: u32, + pub pad1: u32, +} + +/// Scalar parameters for the clipped value-loss gradient (uniform; 32 bytes). +#[repr(C)] +#[derive(Clone, Copy)] +#[cfg_attr( + not(any(target_arch = "spirv", target_arch = "nvptx64")), + derive(bytemuck::Pod, bytemuck::Zeroable) +)] +pub struct PpoValueParams { + /// PPO clip epsilon (value clipping range). + pub clip: f32, + /// Value-loss coefficient. + pub value_coef: f32, + /// Per-sample averaging factor `1 / minibatch_size`. + pub scale: f32, + /// Number of sample columns `M`. + pub num_cols: u32, + pub pad0: u32, + pub pad1: u32, + pub pad2: u32, + pub pad3: u32, +} + +/// Clipped-surrogate actor gradient + log_std gradient contribution, per sample. +/// +/// For sample column `m` (one thread): compute the new diagonal-Gaussian +/// log-prob over the `action_dim` rows, the importance ratio +/// `exp(logp − logp_old)`, the PPO clip mask, then write `g_mean[k,m]` and +/// `g_logstd[k,m]` for every action dim `k`. Matches `minibatch_step`: +/// if !clipped: g_mean = −(adv·ratio·d/σ²)·scale, +/// g_logstd += −adv·ratio·(d²/σ² − 1)·scale, +/// always: g_logstd += −entropy_coef·scale. +#[spirv_bindgen] +#[spirv(compute(threads(256, 1, 1)))] +pub fn gpu_ppo_actor_grad( + #[spirv(global_invocation_id)] invocation_id: UVec3, + #[spirv(uniform, descriptor_set = 0, binding = 0)] params: &PpoActorParams, + #[spirv(storage_buffer, descriptor_set = 0, binding = 1)] mean: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 2)] action: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 3)] log_std: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 4)] adv: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 5)] logp_old: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 6)] g_mean: &mut [f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 7)] g_logstd: &mut [f32], +) { + let a = params.action_dim as usize; + let m_cols = params.num_cols as usize; + let clip = params.clip; + let scale = params.scale; + let ent = params.entropy_coef; + for m in (invocation_id.x as usize..m_cols).step_by(MAX_NUM_THREADS as usize) { + // New log-prob over the action dims (matches ActorCritic::logp). + let mut logp = 0.0f32; + for k in 0..a { + let idx = k * m_cols + m; + let ls = log_std.read(k); + let std = ls.exp(); + let d = (action.read(idx) - mean.read(idx)) / std; + logp += -0.5 * d * d - ls - params.log_sqrt_2pi; + } + let ratio = (logp - logp_old.read(m)).exp(); + let av = adv.read(m); + let clipped = + (av >= 0.0 && ratio > 1.0 + clip) || (av < 0.0 && ratio < 1.0 - clip); + for k in 0..a { + let idx = k * m_cols + m; + let ls = log_std.read(k); + let inv_var = (-2.0 * ls).exp(); // 1/σ² + if clipped { + *g_mean.at_mut(idx) = 0.0; + *g_logstd.at_mut(idx) = -ent * scale; + } else { + let d = action.read(idx) - mean.read(idx); + *g_mean.at_mut(idx) = -(av * ratio * d * inv_var) * scale; + let dls = av * ratio * (d * d * inv_var - 1.0); + *g_logstd.at_mut(idx) = (-dls - ent) * scale; + } + } + } +} + +/// Clipped value-loss gradient, per sample. +/// +/// For sample column `m`: `v_clipped = value_old + clamp(v − value_old, ±clip)`, +/// and `dv = 2·(v_clipped − ret)` if the clipped squared error is larger else +/// `2·(v − ret)`; writes `g_v[m] = value_coef·dv·scale`. Matches `minibatch_step`. +#[spirv_bindgen] +#[spirv(compute(threads(256, 1, 1)))] +pub fn gpu_ppo_value_grad( + #[spirv(global_invocation_id)] invocation_id: UVec3, + #[spirv(uniform, descriptor_set = 0, binding = 0)] params: &PpoValueParams, + #[spirv(storage_buffer, descriptor_set = 0, binding = 1)] v_pred: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 2)] value_old: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 3)] ret: &[f32], + #[spirv(storage_buffer, descriptor_set = 0, binding = 4)] g_v: &mut [f32], +) { + let m_cols = params.num_cols as usize; + let clip = params.clip; + let scale = params.scale; + for m in (invocation_id.x as usize..m_cols).step_by(MAX_NUM_THREADS as usize) { + let v = v_pred.read(m); + let vo = value_old.read(m); + let r = ret.read(m); + let diff = v - vo; + let clamped = if diff > clip { + clip + } else if diff < -clip { + -clip + } else { + diff + }; + let v_clipped = vo + clamped; + let l_un = (v - r) * (v - r); + let l_cl = (v_clipped - r) * (v_clipped - r); + let dv = if l_cl > l_un { + 2.0 * (v_clipped - r) + } else { + 2.0 * (v - r) + }; + *g_v.at_mut(m) = params.value_coef * dv * scale; + } +}