use scirs2_core::gpu::{GpuBackend, GpuBuffer, GpuContext, GpuKernelHandle};
use scirs2_core::ndarray::{Array, Dimension};
use crate::shaders::{OptimizerKernel, WORKGROUP_SIZE};
use crate::{GpuOptimError, GpuOptimizer};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct GpuOptimizerConfig {
pub backend: Option<GpuBackend>,
}
impl GpuOptimizerConfig {
pub fn with_backend(backend: GpuBackend) -> Self {
Self {
backend: Some(backend),
}
}
}
pub const SUPPORTED_BACKENDS: [GpuBackend; 2] = [GpuBackend::Wgpu, GpuBackend::Metal];
fn workgroup_count(n: usize) -> Result<u32, GpuOptimError> {
let groups = n.div_ceil(WORKGROUP_SIZE);
u32::try_from(groups).map_err(|_| {
GpuOptimError::UnsupportedOperation(format!(
"{n} elements need {groups} workgroups, which exceeds the u32 dispatch limit"
))
})
}
fn encode_u32(value: usize) -> Result<f32, GpuOptimError> {
let raw = u32::try_from(value).map_err(|_| {
GpuOptimError::UnsupportedOperation(format!("{value} does not fit in a u32 kernel operand"))
})?;
Ok(f32::from_bits(raw))
}
struct KernelCache {
context: GpuContext,
backend: GpuBackend,
installed: Option<&'static str>,
handle: Option<GpuKernelHandle>,
}
impl KernelCache {
fn new(requested: Option<GpuBackend>) -> Result<Self, GpuOptimError> {
match requested {
Some(backend) => {
if !SUPPORTED_BACKENDS.contains(&backend) {
return Err(GpuOptimError::UnsupportedOperation(format!(
"backend {backend} cannot run optirs-gpu optimizer kernels; \
supported backends are {SUPPORTED_BACKENDS:?}"
)));
}
let context = GpuContext::new(backend)?;
Ok(Self {
context,
backend,
installed: None,
handle: None,
})
}
None => {
let mut reasons = Vec::new();
for backend in SUPPORTED_BACKENDS {
match GpuContext::new(backend) {
Ok(context) => {
return Ok(Self {
context,
backend,
installed: None,
handle: None,
})
}
Err(e) => reasons.push(format!("{backend}: {e}")),
}
}
Err(GpuOptimError::UnsupportedOperation(format!(
"no GPU backend available for optimizer kernels ({})",
reasons.join("; ")
)))
}
}
}
fn context(&self) -> &GpuContext {
&self.context
}
fn backend(&self) -> GpuBackend {
self.backend
}
fn kernel(&mut self, kernel: OptimizerKernel) -> Result<&GpuKernelHandle, GpuOptimError> {
let key = kernel.cache_key(self.backend);
if self.installed != Some(key) || self.handle.is_none() {
let source = kernel.source_for(self.backend).ok_or_else(|| {
GpuOptimError::UnsupportedOperation(format!(
"no {} shader source for backend {}",
kernel.id(),
self.backend
))
})?;
let handle = self.context.execute(|compiler| compiler.compile(source))?;
self.handle = Some(handle);
self.installed = Some(key);
}
self.handle.as_ref().ok_or_else(|| {
GpuOptimError::InvalidState("kernel compilation produced no handle".into())
})
}
}
struct StateBuffers {
slots: usize,
len: usize,
host: Vec<Vec<f32>>,
device: Vec<GpuBuffer<f32>>,
}
impl StateBuffers {
fn new(slots: usize, len: usize) -> Self {
Self {
slots,
len,
host: vec![vec![0.0f32; len]; slots],
device: Vec::new(),
}
}
fn is_resident(&self) -> bool {
self.device.len() == self.slots
}
fn upload(&mut self, context: &GpuContext) -> Result<(), GpuOptimError> {
if self.len == 0 {
return Err(GpuOptimError::InvalidState(
"cannot allocate zero-length optimizer state".into(),
));
}
let mut device = Vec::with_capacity(self.slots);
for slot in &self.host {
let buffer = context.create_buffer::<f32>(self.len);
buffer.copy_from_host(slot)?;
device.push(buffer);
}
self.device = device;
Ok(())
}
fn download(&mut self) -> Result<(), GpuOptimError> {
if !self.is_resident() {
return Ok(());
}
for (slot, buffer) in self.host.iter_mut().zip(self.device.iter()) {
buffer.copy_to_host(slot)?;
}
self.device.clear();
Ok(())
}
fn device_slot(&self, index: usize) -> Result<&GpuBuffer<f32>, GpuOptimError> {
self.device.get(index).ok_or(GpuOptimError::NotInitialized)
}
fn resize(&mut self, len: usize) {
self.len = len;
self.host = vec![vec![0.0f32; len]; self.slots];
self.device.clear();
}
}
struct GpuStepEngine {
cache: KernelCache,
state: StateBuffers,
on_gpu: bool,
step_count: u64,
}
impl GpuStepEngine {
fn new(config: GpuOptimizerConfig, slots: usize) -> Result<Self, GpuOptimError> {
Ok(Self {
cache: KernelCache::new(config.backend)?,
state: StateBuffers::new(slots, 0),
on_gpu: false,
step_count: 0,
})
}
fn backend(&self) -> GpuBackend {
self.cache.backend()
}
fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
if self.on_gpu {
return Ok(());
}
if self.state.len > 0 {
let context = &self.cache.context;
self.state.upload(context)?;
}
self.on_gpu = true;
Ok(())
}
fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
if !self.on_gpu {
return Ok(());
}
self.state.download()?;
self.on_gpu = false;
Ok(())
}
fn prepare(&mut self, len: usize) -> Result<(), GpuOptimError> {
if !self.on_gpu {
return Err(GpuOptimError::InvalidState(
"optimizer is on the CPU; call move_to_gpu() before step_gpu()".into(),
));
}
if len == 0 {
return Err(GpuOptimError::InvalidState(
"cannot run a GPU step on an empty parameter array".into(),
));
}
const MAX_BUFFER_BYTES: usize = 1024 * 1024 * 1024;
let bytes = len.saturating_mul(std::mem::size_of::<f32>());
if bytes > MAX_BUFFER_BYTES {
return Err(GpuOptimError::UnsupportedOperation(format!(
"{len} f32 parameters need {bytes} bytes, above the {MAX_BUFFER_BYTES}-byte \
per-buffer limit of the GPU backends this crate supports"
)));
}
if self.state.len != len {
self.state.resize(len);
self.step_count = 0;
}
if !self.state.is_resident() {
let context = &self.cache.context;
self.state.upload(context)?;
}
Ok(())
}
}
fn to_host_vec<D: Dimension>(array: &Array<f32, D>) -> Vec<f32> {
match array.as_slice() {
Some(slice) => slice.to_vec(),
None => array.iter().copied().collect(),
}
}
fn from_host_vec<D: Dimension>(
array: &mut Array<f32, D>,
values: &[f32],
) -> Result<(), GpuOptimError> {
if values.len() != array.len() {
return Err(GpuOptimError::DimensionMismatch {
expected: array.shape().to_vec(),
actual: vec![values.len()],
});
}
for (dst, src) in array.iter_mut().zip(values.iter()) {
*dst = *src;
}
Ok(())
}
fn check_shapes<D: Dimension>(
params: &Array<f32, D>,
gradients: &Array<f32, D>,
) -> Result<(), GpuOptimError> {
if params.shape() != gradients.shape() {
return Err(GpuOptimError::DimensionMismatch {
expected: params.shape().to_vec(),
actual: gradients.shape().to_vec(),
});
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct AdamParams {
pub learning_rate: f32,
pub beta1: f32,
pub beta2: f32,
pub epsilon: f32,
pub weight_decay: f32,
}
impl Default for AdamParams {
fn default() -> Self {
Self {
learning_rate: 1e-3,
beta1: 0.9,
beta2: 0.999,
epsilon: 1e-8,
weight_decay: 0.0,
}
}
}
impl AdamParams {
pub fn new(
learning_rate: f32,
beta1: f32,
beta2: f32,
epsilon: f32,
weight_decay: f32,
) -> Result<Self, GpuOptimError> {
let params = Self {
learning_rate,
beta1,
beta2,
epsilon,
weight_decay,
};
params.validate()?;
Ok(params)
}
fn validate(&self) -> Result<(), GpuOptimError> {
let invalid = |what: &str| GpuOptimError::InvalidState(format!("invalid Adam {what}"));
if !(self.learning_rate.is_finite() && self.learning_rate > 0.0) {
return Err(invalid("learning rate (must be finite and > 0)"));
}
if !(self.beta1.is_finite() && (0.0..1.0).contains(&self.beta1)) {
return Err(invalid("beta1 (must be in [0, 1))"));
}
if !(self.beta2.is_finite() && (0.0..1.0).contains(&self.beta2)) {
return Err(invalid("beta2 (must be in [0, 1))"));
}
if !(self.epsilon.is_finite() && self.epsilon > 0.0) {
return Err(invalid("epsilon (must be finite and > 0)"));
}
if !(self.weight_decay.is_finite() && self.weight_decay >= 0.0) {
return Err(invalid("weight decay (must be finite and >= 0)"));
}
Ok(())
}
fn bias_corrections(&self, step: u64) -> (f32, f32) {
let exp = step.min(i32::MAX as u64) as i32;
(1.0 - self.beta1.powi(exp), 1.0 - self.beta2.powi(exp))
}
}
macro_rules! adam_family {
($name:ident, $kernel:expr, $doc:literal) => {
#[doc = $doc]
pub struct $name {
engine: GpuStepEngine,
params: AdamParams,
}
impl $name {
pub fn new(params: AdamParams) -> Result<Self, GpuOptimError> {
Self::with_config(params, GpuOptimizerConfig::default())
}
pub fn with_config(
params: AdamParams,
config: GpuOptimizerConfig,
) -> Result<Self, GpuOptimError> {
params.validate()?;
Ok(Self {
engine: GpuStepEngine::new(config, 2)?,
params,
})
}
pub fn backend(&self) -> GpuBackend {
self.engine.backend()
}
pub fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
pub fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_gpu`")]
pub fn to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_gpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_cpu`")]
pub fn to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_cpu()
}
pub fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
pub fn step_count(&self) -> u64 {
self.engine.step_count
}
pub fn params(&self) -> &AdamParams {
&self.params
}
pub fn set_params(&mut self, params: AdamParams) -> Result<(), GpuOptimError> {
params.validate()?;
self.params = params;
Ok(())
}
}
impl<D: Dimension> GpuOptimizer<f32, D> for $name {
fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
fn step_gpu(
&mut self,
params: &mut Array<f32, D>,
gradients: &Array<f32, D>,
) -> Result<(), GpuOptimError> {
check_shapes(params, gradients)?;
let n = params.len();
self.engine.prepare(n)?;
self.engine.step_count = self.engine.step_count.saturating_add(1);
let (bc1, bc2) = self.params.bias_corrections(self.engine.step_count);
let hyper = [
self.params.learning_rate,
self.params.beta1,
self.params.beta2,
self.params.epsilon,
self.params.weight_decay,
bc1,
bc2,
encode_u32(n)?,
];
let host_params = to_host_vec(params);
let host_grads = to_host_vec(gradients);
let groups = workgroup_count(n)?;
let updated = {
let context = self.engine.cache.context();
let params_buf = context.create_buffer::<f32>(n);
params_buf.copy_from_host(&host_params)?;
let grads_buf = context.create_buffer::<f32>(n);
grads_buf.copy_from_host(&host_grads)?;
let hyper_buf = context.create_buffer::<f32>(hyper.len());
hyper_buf.copy_from_host(&hyper)?;
let m_buf = self.engine.state.device_slot(0)?.clone();
let v_buf = self.engine.state.device_slot(1)?.clone();
let kernel = self.engine.cache.kernel($kernel)?;
kernel.set_buffer("x", ¶ms_buf);
kernel.set_buffer("y", &grads_buf);
kernel.set_buffer("a", &m_buf);
kernel.set_buffer("b", &v_buf);
kernel.set_buffer("result", &hyper_buf);
kernel.dispatch([groups, 1, 1]);
let mut out = vec![0.0f32; n];
params_buf.copy_to_host(&mut out)?;
out
};
from_host_vec(params, &updated)
}
}
};
}
adam_family!(
GpuAdam,
OptimizerKernel::Adam,
"GPU Adam with coupled L2 weight decay, numerically matching \
`optirs_core::optimizers::Adam`."
);
adam_family!(
GpuAdamW,
OptimizerKernel::AdamW,
"GPU AdamW with *decoupled* weight decay: the decay term is applied to the \
parameter and never enters the moment estimates."
);
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct SgdParams {
pub learning_rate: f32,
pub momentum: f32,
pub dampening: f32,
pub weight_decay: f32,
pub nesterov: bool,
}
impl Default for SgdParams {
fn default() -> Self {
Self {
learning_rate: 1e-2,
momentum: 0.0,
dampening: 0.0,
weight_decay: 0.0,
nesterov: false,
}
}
}
impl SgdParams {
fn validate(&self) -> Result<(), GpuOptimError> {
let invalid = |what: &str| GpuOptimError::InvalidState(format!("invalid SGD {what}"));
if !(self.learning_rate.is_finite() && self.learning_rate > 0.0) {
return Err(invalid("learning rate (must be finite and > 0)"));
}
if !(self.momentum.is_finite() && self.momentum >= 0.0) {
return Err(invalid("momentum (must be finite and >= 0)"));
}
if !(self.dampening.is_finite() && (0.0..1.0).contains(&self.dampening)) {
return Err(invalid("dampening (must be in [0, 1))"));
}
if !(self.weight_decay.is_finite() && self.weight_decay >= 0.0) {
return Err(invalid("weight decay (must be finite and >= 0)"));
}
if self.nesterov && (self.momentum <= 0.0 || self.dampening != 0.0) {
return Err(invalid(
"Nesterov mode (requires momentum > 0 and dampening == 0)",
));
}
Ok(())
}
}
pub struct GpuSgd {
engine: GpuStepEngine,
params: SgdParams,
}
impl GpuSgd {
pub fn new(params: SgdParams) -> Result<Self, GpuOptimError> {
Self::with_config(params, GpuOptimizerConfig::default())
}
pub fn with_config(
params: SgdParams,
config: GpuOptimizerConfig,
) -> Result<Self, GpuOptimError> {
params.validate()?;
Ok(Self {
engine: GpuStepEngine::new(config, 1)?,
params,
})
}
pub fn backend(&self) -> GpuBackend {
self.engine.backend()
}
pub fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
pub fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_gpu`")]
pub fn to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_gpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_cpu`")]
pub fn to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_cpu()
}
pub fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
pub fn step_count(&self) -> u64 {
self.engine.step_count
}
}
impl<D: Dimension> GpuOptimizer<f32, D> for GpuSgd {
fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
fn step_gpu(
&mut self,
params: &mut Array<f32, D>,
gradients: &Array<f32, D>,
) -> Result<(), GpuOptimError> {
check_shapes(params, gradients)?;
let n = params.len();
self.engine.prepare(n)?;
let first = self.engine.step_count == 0;
self.engine.step_count = self.engine.step_count.saturating_add(1);
let hyper = [
self.params.learning_rate,
self.params.momentum,
self.params.dampening,
self.params.weight_decay,
if self.params.nesterov { 1.0 } else { 0.0 },
if first { 1.0 } else { 0.0 },
encode_u32(n)?,
];
let host_params = to_host_vec(params);
let host_grads = to_host_vec(gradients);
let groups = workgroup_count(n)?;
let updated = {
let context = self.engine.cache.context();
let params_buf = context.create_buffer::<f32>(n);
params_buf.copy_from_host(&host_params)?;
let grads_buf = context.create_buffer::<f32>(n);
grads_buf.copy_from_host(&host_grads)?;
let hyper_buf = context.create_buffer::<f32>(hyper.len());
hyper_buf.copy_from_host(&hyper)?;
let buf = self.engine.state.device_slot(0)?.clone();
let kernel = self.engine.cache.kernel(OptimizerKernel::Sgd)?;
kernel.set_buffer("x", ¶ms_buf);
kernel.set_buffer("y", &grads_buf);
kernel.set_buffer("a", &buf);
kernel.set_buffer("b", &hyper_buf);
kernel.dispatch([groups, 1, 1]);
let mut out = vec![0.0f32; n];
params_buf.copy_to_host(&mut out)?;
out
};
from_host_vec(params, &updated)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct RmspropParams {
pub learning_rate: f32,
pub alpha: f32,
pub epsilon: f32,
pub weight_decay: f32,
pub momentum: f32,
pub centered: bool,
}
impl Default for RmspropParams {
fn default() -> Self {
Self {
learning_rate: 1e-2,
alpha: 0.99,
epsilon: 1e-8,
weight_decay: 0.0,
momentum: 0.0,
centered: false,
}
}
}
impl RmspropParams {
fn validate(&self) -> Result<(), GpuOptimError> {
let invalid = |what: &str| GpuOptimError::InvalidState(format!("invalid RMSprop {what}"));
if !(self.learning_rate.is_finite() && self.learning_rate > 0.0) {
return Err(invalid("learning rate (must be finite and > 0)"));
}
if !(self.alpha.is_finite() && (0.0..1.0).contains(&self.alpha)) {
return Err(invalid("alpha (must be in [0, 1))"));
}
if !(self.epsilon.is_finite() && self.epsilon > 0.0) {
return Err(invalid("epsilon (must be finite and > 0)"));
}
if !(self.weight_decay.is_finite() && self.weight_decay >= 0.0) {
return Err(invalid("weight decay (must be finite and >= 0)"));
}
if !(self.momentum.is_finite() && self.momentum >= 0.0) {
return Err(invalid("momentum (must be finite and >= 0)"));
}
Ok(())
}
}
pub struct GpuRmsprop {
engine: GpuStepEngine,
params: RmspropParams,
}
impl GpuRmsprop {
pub fn new(params: RmspropParams) -> Result<Self, GpuOptimError> {
Self::with_config(params, GpuOptimizerConfig::default())
}
pub fn with_config(
params: RmspropParams,
config: GpuOptimizerConfig,
) -> Result<Self, GpuOptimError> {
params.validate()?;
Ok(Self {
engine: GpuStepEngine::new(config, 3)?,
params,
})
}
pub fn backend(&self) -> GpuBackend {
self.engine.backend()
}
pub fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
pub fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_gpu`")]
pub fn to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_gpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_cpu`")]
pub fn to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_cpu()
}
pub fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
pub fn step_count(&self) -> u64 {
self.engine.step_count
}
}
impl<D: Dimension> GpuOptimizer<f32, D> for GpuRmsprop {
fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
fn step_gpu(
&mut self,
params: &mut Array<f32, D>,
gradients: &Array<f32, D>,
) -> Result<(), GpuOptimError> {
check_shapes(params, gradients)?;
let n = params.len();
self.engine.prepare(n)?;
self.engine.step_count = self.engine.step_count.saturating_add(1);
let hyper = [
self.params.learning_rate,
self.params.alpha,
self.params.epsilon,
self.params.weight_decay,
self.params.momentum,
if self.params.centered { 1.0 } else { 0.0 },
encode_u32(n)?,
];
let host_params = to_host_vec(params);
let host_grads = to_host_vec(gradients);
let groups = workgroup_count(n)?;
let updated = {
let context = self.engine.cache.context();
let params_buf = context.create_buffer::<f32>(n);
params_buf.copy_from_host(&host_params)?;
let grads_buf = context.create_buffer::<f32>(n);
grads_buf.copy_from_host(&host_grads)?;
let hyper_buf = context.create_buffer::<f32>(hyper.len());
hyper_buf.copy_from_host(&hyper)?;
let sq_avg = self.engine.state.device_slot(0)?.clone();
let g_avg = self.engine.state.device_slot(1)?.clone();
let buf = self.engine.state.device_slot(2)?.clone();
let kernel = self.engine.cache.kernel(OptimizerKernel::Rmsprop)?;
kernel.set_buffer("x", ¶ms_buf);
kernel.set_buffer("y", &grads_buf);
kernel.set_buffer("a", &sq_avg);
kernel.set_buffer("b", &g_avg);
kernel.set_buffer("result", &buf);
kernel.set_buffer("output", &hyper_buf);
kernel.dispatch([groups, 1, 1]);
let mut out = vec![0.0f32; n];
params_buf.copy_to_host(&mut out)?;
out
};
from_host_vec(params, &updated)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct AdagradParams {
pub learning_rate: f32,
pub lr_decay: f32,
pub epsilon: f32,
pub weight_decay: f32,
}
impl Default for AdagradParams {
fn default() -> Self {
Self {
learning_rate: 1e-2,
lr_decay: 0.0,
epsilon: 1e-10,
weight_decay: 0.0,
}
}
}
impl AdagradParams {
fn validate(&self) -> Result<(), GpuOptimError> {
let invalid = |what: &str| GpuOptimError::InvalidState(format!("invalid Adagrad {what}"));
if !(self.learning_rate.is_finite() && self.learning_rate > 0.0) {
return Err(invalid("learning rate (must be finite and > 0)"));
}
if !(self.lr_decay.is_finite() && self.lr_decay >= 0.0) {
return Err(invalid("lr_decay (must be finite and >= 0)"));
}
if !(self.epsilon.is_finite() && self.epsilon > 0.0) {
return Err(invalid("epsilon (must be finite and > 0)"));
}
if !(self.weight_decay.is_finite() && self.weight_decay >= 0.0) {
return Err(invalid("weight decay (must be finite and >= 0)"));
}
Ok(())
}
fn effective_lr(&self, step: u64) -> f32 {
let completed = step.saturating_sub(1) as f32;
self.learning_rate / (1.0 + completed * self.lr_decay)
}
}
pub struct GpuAdagrad {
engine: GpuStepEngine,
params: AdagradParams,
}
impl GpuAdagrad {
pub fn new(params: AdagradParams) -> Result<Self, GpuOptimError> {
Self::with_config(params, GpuOptimizerConfig::default())
}
pub fn with_config(
params: AdagradParams,
config: GpuOptimizerConfig,
) -> Result<Self, GpuOptimError> {
params.validate()?;
Ok(Self {
engine: GpuStepEngine::new(config, 1)?,
params,
})
}
pub fn backend(&self) -> GpuBackend {
self.engine.backend()
}
pub fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
pub fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_gpu`")]
pub fn to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_gpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_cpu`")]
pub fn to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_cpu()
}
pub fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
pub fn step_count(&self) -> u64 {
self.engine.step_count
}
pub fn next_effective_lr(&self) -> f32 {
self.params.effective_lr(self.engine.step_count + 1)
}
}
impl<D: Dimension> GpuOptimizer<f32, D> for GpuAdagrad {
fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
fn step_gpu(
&mut self,
params: &mut Array<f32, D>,
gradients: &Array<f32, D>,
) -> Result<(), GpuOptimError> {
check_shapes(params, gradients)?;
let n = params.len();
self.engine.prepare(n)?;
self.engine.step_count = self.engine.step_count.saturating_add(1);
let hyper = [
self.params.effective_lr(self.engine.step_count),
self.params.epsilon,
self.params.weight_decay,
encode_u32(n)?,
];
let host_params = to_host_vec(params);
let host_grads = to_host_vec(gradients);
let groups = workgroup_count(n)?;
let updated = {
let context = self.engine.cache.context();
let params_buf = context.create_buffer::<f32>(n);
params_buf.copy_from_host(&host_params)?;
let grads_buf = context.create_buffer::<f32>(n);
grads_buf.copy_from_host(&host_grads)?;
let hyper_buf = context.create_buffer::<f32>(hyper.len());
hyper_buf.copy_from_host(&hyper)?;
let sum = self.engine.state.device_slot(0)?.clone();
let kernel = self.engine.cache.kernel(OptimizerKernel::Adagrad)?;
kernel.set_buffer("x", ¶ms_buf);
kernel.set_buffer("y", &grads_buf);
kernel.set_buffer("a", &sum);
kernel.set_buffer("b", &hyper_buf);
kernel.dispatch([groups, 1, 1]);
let mut out = vec![0.0f32; n];
params_buf.copy_to_host(&mut out)?;
out
};
from_host_vec(params, &updated)
}
}
pub struct GpuLamb {
engine: GpuStepEngine,
params: AdamParams,
}
impl GpuLamb {
pub fn new(params: AdamParams) -> Result<Self, GpuOptimError> {
Self::with_config(params, GpuOptimizerConfig::default())
}
pub fn with_config(
params: AdamParams,
config: GpuOptimizerConfig,
) -> Result<Self, GpuOptimError> {
params.validate()?;
Ok(Self {
engine: GpuStepEngine::new(config, 2)?,
params,
})
}
pub fn backend(&self) -> GpuBackend {
self.engine.backend()
}
pub fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
pub fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_gpu`")]
pub fn to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_gpu()
}
#[deprecated(since = "0.3.2", note = "renamed to `move_to_cpu`")]
pub fn to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.move_to_cpu()
}
pub fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
pub fn step_count(&self) -> u64 {
self.engine.step_count
}
fn trust_ratio(param_norm: f32, update_norm: f32) -> f32 {
if param_norm > 0.0 && update_norm > 0.0 {
param_norm / update_norm
} else {
1.0
}
}
}
impl<D: Dimension> GpuOptimizer<f32, D> for GpuLamb {
fn is_gpu_available(&self) -> bool {
self.engine.backend() != GpuBackend::Cpu
}
fn move_to_gpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_gpu()
}
fn move_to_cpu(&mut self) -> Result<(), GpuOptimError> {
self.engine.move_to_cpu()
}
fn step_gpu(
&mut self,
params: &mut Array<f32, D>,
gradients: &Array<f32, D>,
) -> Result<(), GpuOptimError> {
check_shapes(params, gradients)?;
let n = params.len();
self.engine.prepare(n)?;
self.engine.step_count = self.engine.step_count.saturating_add(1);
let (bc1, bc2) = self.params.bias_corrections(self.engine.step_count);
let groups = workgroup_count(n)?;
let partial_len = (groups as usize).saturating_mul(2);
let host_params = to_host_vec(params);
let host_grads = to_host_vec(gradients);
let mut hyper = [
self.params.learning_rate,
self.params.beta1,
self.params.beta2,
self.params.epsilon,
self.params.weight_decay,
bc1,
bc2,
encode_u32(n)?,
encode_u32(0)?,
1.0,
];
let updated = {
let context = self.engine.cache.context();
let params_buf = context.create_buffer::<f32>(n);
params_buf.copy_from_host(&host_params)?;
let grads_buf = context.create_buffer::<f32>(n);
grads_buf.copy_from_host(&host_grads)?;
let scratch_len = n.saturating_add(partial_len);
let scratch_buf = context.create_buffer::<f32>(scratch_len);
scratch_buf.copy_from_host(&vec![0.0f32; scratch_len])?;
let hyper_buf = context.create_buffer::<f32>(hyper.len());
hyper_buf.copy_from_host(&hyper)?;
let m_buf = self.engine.state.device_slot(0)?.clone();
let v_buf = self.engine.state.device_slot(1)?.clone();
{
let kernel = self.engine.cache.kernel(OptimizerKernel::Lamb)?;
kernel.set_buffer("x", ¶ms_buf);
kernel.set_buffer("y", &grads_buf);
kernel.set_buffer("a", &m_buf);
kernel.set_buffer("b", &v_buf);
kernel.set_buffer("result", &scratch_buf);
kernel.set_buffer("output", &hyper_buf);
kernel.dispatch([groups, 1, 1]);
}
let mut scratch = vec![0.0f32; scratch_len];
scratch_buf.copy_to_host(&mut scratch)?;
let mut param_sq = 0.0f64;
let mut update_sq = 0.0f64;
for pair in scratch[n..].chunks_exact(2) {
param_sq += f64::from(pair[0]);
update_sq += f64::from(pair[1]);
}
let trust = Self::trust_ratio(
param_sq.max(0.0).sqrt() as f32,
update_sq.max(0.0).sqrt() as f32,
);
hyper[8] = encode_u32(1)?;
hyper[9] = trust;
hyper_buf.copy_from_host(&hyper)?;
{
let kernel = self.engine.cache.kernel(OptimizerKernel::Lamb)?;
kernel.set_buffer("x", ¶ms_buf);
kernel.set_buffer("y", &grads_buf);
kernel.set_buffer("a", &m_buf);
kernel.set_buffer("b", &v_buf);
kernel.set_buffer("result", &scratch_buf);
kernel.set_buffer("output", &hyper_buf);
kernel.dispatch([groups, 1, 1]);
}
let mut out = vec![0.0f32; n];
params_buf.copy_to_host(&mut out)?;
out
};
from_host_vec(params, &updated)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn adam_params_reject_invalid_values() {
assert!(AdamParams::new(1e-3, 0.9, 0.999, 1e-8, 0.0).is_ok());
assert!(AdamParams::new(0.0, 0.9, 0.999, 1e-8, 0.0).is_err());
assert!(AdamParams::new(1e-3, 1.0, 0.999, 1e-8, 0.0).is_err());
assert!(AdamParams::new(1e-3, 0.9, f32::NAN, 1e-8, 0.0).is_err());
assert!(AdamParams::new(1e-3, 0.9, 0.999, 0.0, 0.0).is_err());
assert!(AdamParams::new(1e-3, 0.9, 0.999, 1e-8, -1.0).is_err());
}
#[test]
fn adagrad_lr_decay_is_not_off_by_one() {
let params = AdagradParams {
learning_rate: 0.1,
lr_decay: 0.5,
..AdagradParams::default()
};
assert!((params.effective_lr(1) - 0.1).abs() < 1e-9);
assert!((params.effective_lr(2) - 0.1 / 1.5).abs() < 1e-9);
assert!((params.effective_lr(3) - 0.1 / 2.0).abs() < 1e-9);
}
#[test]
fn encode_u32_round_trips_large_counts() {
let n = 20_000_000usize; let encoded = encode_u32(n).expect("encodable");
assert_eq!(encoded.to_bits() as usize, n);
}
#[test]
fn workgroup_count_covers_the_tail() {
assert_eq!(workgroup_count(1).expect("ok"), 1);
assert_eq!(workgroup_count(256).expect("ok"), 1);
assert_eq!(workgroup_count(257).expect("ok"), 2);
assert_eq!(workgroup_count(0).expect("ok"), 0);
}
#[test]
fn lamb_trust_ratio_degenerates_to_one() {
assert_eq!(GpuLamb::trust_ratio(0.0, 1.0), 1.0);
assert_eq!(GpuLamb::trust_ratio(1.0, 0.0), 1.0);
assert!((GpuLamb::trust_ratio(4.0, 2.0) - 2.0).abs() < 1e-6);
}
#[test]
fn state_buffers_start_zeroed_on_the_host() {
let state = StateBuffers::new(2, 8);
assert!(!state.is_resident());
assert_eq!(state.host.len(), 2);
assert!(state.host.iter().all(|s| s.iter().all(|&x| x == 0.0)));
}
#[test]
fn non_wgpu_backends_are_rejected_explicitly() {
for backend in [
GpuBackend::Cpu,
GpuBackend::Cuda,
GpuBackend::Rocm,
GpuBackend::OpenCL,
] {
let err = KernelCache::new(Some(backend))
.err()
.unwrap_or_else(|| panic!("backend {backend} must be rejected"));
assert!(
matches!(err, GpuOptimError::UnsupportedOperation(_)),
"backend {backend} produced the wrong error: {err}"
);
}
}
}