use super::*;
use crate::webgpu::backend::ComputeBackend;
use crate::webgpu::{ComputeContext, ComputeError};
use num_traits::Float;
use std::collections::HashMap;
use std::sync::Arc;
#[cfg(feature = "gpu")]
pub struct GpuAdam<T: Float + Send + Sync + Default + std::fmt::Debug + 'static> {
learning_rate: T,
beta1: T,
beta2: T,
epsilon: T,
weight_decay: T,
error_function: Box<dyn ErrorFunction<T>>,
_compute_context: ComputeContext<T>,
webgpu_backend: Option<Arc<dyn ComputeBackend<T>>>,
m_weights_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
v_weights_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
m_biases_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
v_biases_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
step: usize,
gpu_stats: GpuPerformanceStats,
moment_estimates: Option<HashMap<String, T>>,
callback: Option<TrainingCallback<T>>,
}
#[cfg(feature = "gpu")]
#[allow(dead_code)]
pub struct GpuAdamW<T: Float + Send + Sync + Default + std::fmt::Debug + 'static> {
learning_rate: T,
beta1: T,
beta2: T,
epsilon: T,
weight_decay: T,
error_function: Box<dyn ErrorFunction<T>>,
_compute_context: ComputeContext<T>,
webgpu_backend: Option<Arc<dyn ComputeBackend<T>>>,
m_weights_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
v_weights_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
m_biases_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
v_biases_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
step: usize,
gpu_stats: GpuPerformanceStats,
callback: Option<TrainingCallback<T>>,
}
#[cfg(feature = "gpu")]
#[allow(dead_code)]
pub struct GpuBatchBackprop<T: Float + Send + Sync + Default + std::fmt::Debug + 'static> {
learning_rate: T,
momentum: T,
error_function: Box<dyn ErrorFunction<T>>,
_compute_context: ComputeContext<T>,
webgpu_backend: Option<Arc<dyn ComputeBackend<T>>>,
momentum_weights_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
momentum_biases_gpu: Option<Vec<crate::webgpu::memory::BufferHandle>>,
gpu_stats: GpuPerformanceStats,
callback: Option<TrainingCallback<T>>,
}
#[derive(Debug, Default, Clone)]
pub struct GpuPerformanceStats {
pub total_gpu_time_ms: f64,
pub memory_transfer_time_ms: f64,
pub kernel_launches: u64,
pub avg_batch_time_ms: f64,
pub gpu_memory_used_bytes: u64,
pub speedup_vs_cpu: f64,
}
#[cfg(feature = "gpu")]
impl<T: Float + Send + Sync + Default + std::fmt::Debug + 'static> GpuAdam<T> {
pub fn new(learning_rate: T) -> Result<Self, ComputeError> {
let compute_context = ComputeContext::new()?;
let webgpu_backend = Self::initialize_webgpu_backend()?;
Ok(Self {
learning_rate,
beta1: T::from(0.9).unwrap(),
beta2: T::from(0.999).unwrap(),
epsilon: T::from(1e-8).unwrap(),
weight_decay: T::zero(),
error_function: Box::new(MseError),
_compute_context: compute_context,
webgpu_backend,
m_weights_gpu: None,
v_weights_gpu: None,
m_biases_gpu: None,
v_biases_gpu: None,
step: 0,
gpu_stats: GpuPerformanceStats::default(),
moment_estimates: None,
callback: None,
})
}
fn initialize_webgpu_backend() -> Result<Option<Arc<dyn ComputeBackend<T>>>, ComputeError> {
match crate::webgpu::WebGPUBackend::<T>::new() {
Ok(backend) => {
log::info!("GPU acceleration enabled via WebGPU");
Ok(Some(Arc::new(backend) as Arc<dyn ComputeBackend<T>>))
}
Err(e) => {
log::warn!("GPU initialization failed: {e}, falling back to CPU");
Ok(None)
}
}
}
pub fn is_gpu_available(&self) -> bool {
self.webgpu_backend.is_some()
}
pub fn with_beta1(mut self, beta1: T) -> Self {
self.beta1 = beta1;
self
}
pub fn with_beta2(mut self, beta2: T) -> Self {
self.beta2 = beta2;
self
}
pub fn with_epsilon(mut self, epsilon: T) -> Self {
self.epsilon = epsilon;
self
}
pub fn with_weight_decay(mut self, weight_decay: T) -> Self {
self.weight_decay = weight_decay;
self
}
pub fn with_error_function(mut self, error_function: Box<dyn ErrorFunction<T>>) -> Self {
self.error_function = error_function;
self
}
pub fn get_performance_stats(&self) -> &GpuPerformanceStats {
&self.gpu_stats
}
fn initialize_gpu_buffers(&mut self, network: &Network<T>) -> Result<(), ComputeError> {
if self.m_weights_gpu.is_none() && self.webgpu_backend.is_some() {
let backend = self.webgpu_backend.clone().unwrap();
let memory_manager = backend.memory_manager();
let mut m_weights = Vec::new();
let mut v_weights = Vec::new();
for layer in network.layers.iter().skip(1) {
let num_weights = layer
.neurons
.iter()
.filter(|n| !n.is_bias)
.map(|n| n.connections.len())
.sum::<usize>();
if num_weights > 0 {
let m_buffer =
memory_manager.allocate_buffer(num_weights * std::mem::size_of::<T>())?;
let v_buffer =
memory_manager.allocate_buffer(num_weights * std::mem::size_of::<T>())?;
let zeros = vec![T::zero(); num_weights];
memory_manager.upload_data(m_buffer, &zeros)?;
memory_manager.upload_data(v_buffer, &zeros)?;
m_weights.push(m_buffer);
v_weights.push(v_buffer);
}
}
let mut m_biases = Vec::new();
let mut v_biases = Vec::new();
for layer in network.layers.iter().skip(1) {
let num_biases = layer.neurons.iter().filter(|n| !n.is_bias).count();
if num_biases > 0 {
let m_buffer =
memory_manager.allocate_buffer(num_biases * std::mem::size_of::<T>())?;
let v_buffer =
memory_manager.allocate_buffer(num_biases * std::mem::size_of::<T>())?;
let zeros = vec![T::zero(); num_biases];
memory_manager.upload_data(m_buffer, &zeros)?;
memory_manager.upload_data(v_buffer, &zeros)?;
m_biases.push(m_buffer);
v_biases.push(v_buffer);
}
}
self.m_weights_gpu = Some(m_weights);
self.v_weights_gpu = Some(v_weights);
self.m_biases_gpu = Some(m_biases);
self.v_biases_gpu = Some(v_biases);
}
Ok(())
}
fn gpu_train_step(
&mut self,
network: &mut Network<T>,
data: &TrainingData<T>,
) -> Result<T, ComputeError> {
let start_time = std::time::Instant::now();
self.initialize_gpu_buffers(network)?;
self.step += 1;
if let Some(backend) = self.webgpu_backend.clone() {
use super::gpu_batch_training::gpu_batch_train_step;
let total_error = gpu_batch_train_step(network, data, backend.clone(), self)?;
let elapsed = start_time.elapsed();
self.gpu_stats.total_gpu_time_ms += elapsed.as_secs_f64() * 1000.0;
self.gpu_stats.kernel_launches += 1; self.gpu_stats.avg_batch_time_ms = elapsed.as_secs_f64() * 1000.0;
Ok(total_error)
} else {
Err(ComputeError::GpuUnavailable)
}
}
fn extract_layer_weights(&self, layer: &crate::Layer<T>) -> Vec<T> {
let mut weights = Vec::new();
for neuron in layer.neurons.iter().filter(|n| !n.is_bias) {
for connection in neuron.connections.iter().skip(1) {
weights.push(connection.weight);
}
}
weights
}
fn extract_layer_biases(&self, layer: &crate::Layer<T>) -> Vec<T> {
let mut biases = Vec::new();
for neuron in layer.neurons.iter().filter(|n| !n.is_bias) {
let bias = if !neuron.connections.is_empty() {
neuron.connections[0].weight
} else {
T::zero()
};
biases.push(bias);
}
biases
}
pub(super) fn apply_adam_updates_with_gradients(
&mut self,
network: &mut Network<T>,
weight_gradients: &[Vec<T>],
bias_gradients: &[Vec<T>],
) -> Result<(), ComputeError> {
if self.m_weights_gpu.is_none() {
self.initialize_moment_estimates(network)?;
}
let lr_t = self.learning_rate * (T::one() - self.beta2.powi(self.step as i32)).sqrt()
/ (T::one() - self.beta1.powi(self.step as i32));
for (layer_idx, layer) in network.layers.iter_mut().skip(1).enumerate() {
let weight_grads = &weight_gradients[layer_idx];
let bias_grads = &bias_gradients[layer_idx];
let mut weight_idx = 0;
for (neuron_idx, neuron) in layer.neurons.iter_mut().filter(|n| !n.is_bias).enumerate()
{
if !neuron.connections.is_empty() {
let bias_grad = bias_grads[neuron_idx];
self.update_adam_parameter(
&mut neuron.connections[0].weight,
bias_grad,
lr_t,
layer_idx,
neuron_idx,
true, );
}
for connection in neuron.connections.iter_mut().skip(1) {
if weight_idx < weight_grads.len() {
let weight_grad = weight_grads[weight_idx];
self.update_adam_parameter(
&mut connection.weight,
weight_grad,
lr_t,
layer_idx,
weight_idx,
false, );
weight_idx += 1;
}
}
}
}
Ok(())
}
fn update_adam_parameter(
&mut self,
param: &mut T,
gradient: T,
lr_t: T,
layer_idx: usize,
param_idx: usize,
is_bias: bool,
) {
let m_key = format!("{}_{}_{}_{}", layer_idx, param_idx, is_bias, "m");
let v_key = format!("{}_{}_{}_{}", layer_idx, param_idx, is_bias, "v");
let m = self.get_moment(&m_key).unwrap_or(T::zero());
let v = self.get_moment(&v_key).unwrap_or(T::zero());
let new_m = self.beta1 * m + (T::one() - self.beta1) * gradient;
let new_v = self.beta2 * v + (T::one() - self.beta2) * gradient * gradient;
self.set_moment(&m_key, new_m);
self.set_moment(&v_key, new_v);
*param = *param - lr_t * new_m / (new_v.sqrt() + self.epsilon);
if self.weight_decay > T::zero() && !is_bias {
*param = *param - self.learning_rate * self.weight_decay * *param;
}
}
fn initialize_moment_estimates(&mut self, _network: &Network<T>) -> Result<(), ComputeError> {
self.moment_estimates = Some(HashMap::new());
Ok(())
}
fn get_moment(&self, key: &str) -> Option<T> {
self.moment_estimates.as_ref()?.get(key).copied()
}
fn set_moment(&mut self, key: &str, value: T) {
if let Some(moments) = self.moment_estimates.as_mut() {
moments.insert(key.to_string(), value);
}
}
}
#[cfg(feature = "gpu")]
impl<T: Float + Send + Sync + Default + std::fmt::Debug + 'static> TrainingAlgorithm<T>
for GpuAdam<T>
{
fn train_epoch(
&mut self,
network: &mut Network<T>,
data: &TrainingData<T>,
) -> Result<T, TrainingError> {
match self.gpu_train_step(network, data) {
Ok(error) => Ok(error),
Err(ComputeError::GpuUnavailable) => {
let mut cpu_adam = super::Adam::new(self.learning_rate)
.with_beta1(self.beta1)
.with_beta2(self.beta2)
.with_epsilon(self.epsilon)
.with_weight_decay(self.weight_decay);
println!("GPU not available, falling back to CPU Adam");
cpu_adam.train_epoch(network, data)
}
Err(e) => Err(TrainingError::TrainingFailed(format!(
"GPU training failed: {e}"
))),
}
}
fn calculate_error(&self, network: &Network<T>, data: &TrainingData<T>) -> T {
if let Some(backend) = self.webgpu_backend.clone() {
let mut total_error = T::zero();
for (input, desired_output) in data.inputs.iter().zip(data.outputs.iter()) {
let mut current_input = input.clone();
for layer in network.layers.iter().skip(1) {
let weights = self.extract_layer_weights(layer);
let biases = self.extract_layer_biases(layer);
if let Ok(layer_output) = backend.matrix_vector_multiply(
&weights,
¤t_input,
biases.len(),
current_input.len(),
) {
let mut activated_output = Vec::new();
for (i, &output) in layer_output.iter().enumerate() {
let with_bias = output + biases[i];
activated_output.push(with_bias);
}
let activation_function = layer
.neurons
.iter()
.find(|n| !n.is_bias)
.map(|n| n.activation_function)
.unwrap_or(crate::ActivationFunction::Sigmoid);
if let Ok(activated) = backend.apply_activation_function(
&activated_output,
activation_function,
T::one(),
) {
current_input = activated;
}
}
}
total_error = total_error
+ self
.error_function
.calculate(¤t_input, desired_output);
}
total_error / T::from(data.inputs.len()).unwrap()
} else {
let mut total_error = T::zero();
let mut network_clone = network.clone();
for (input, desired_output) in data.inputs.iter().zip(data.outputs.iter()) {
let output = network_clone.run(input);
total_error = total_error + self.error_function.calculate(&output, desired_output);
}
total_error / T::from(data.inputs.len()).unwrap()
}
}
fn count_bit_fails(
&self,
network: &Network<T>,
data: &TrainingData<T>,
bit_fail_limit: T,
) -> usize {
let mut bit_fails = 0;
let mut network_clone = network.clone();
for (input, desired_output) in data.inputs.iter().zip(data.outputs.iter()) {
let output = network_clone.run(input);
for (&actual, &desired) in output.iter().zip(desired_output.iter()) {
if (actual - desired).abs() > bit_fail_limit {
bit_fails += 1;
}
}
}
bit_fails
}
fn save_state(&self) -> TrainingState<T> {
let mut state = HashMap::new();
state.insert("learning_rate".to_string(), vec![self.learning_rate]);
state.insert("beta1".to_string(), vec![self.beta1]);
state.insert("beta2".to_string(), vec![self.beta2]);
state.insert("epsilon".to_string(), vec![self.epsilon]);
state.insert("weight_decay".to_string(), vec![self.weight_decay]);
state.insert("step".to_string(), vec![T::from(self.step).unwrap()]);
TrainingState {
epoch: 0,
best_error: T::from(f32::MAX).unwrap(),
algorithm_specific: state,
}
}
fn restore_state(&mut self, state: TrainingState<T>) {
if let Some(lr) = state.algorithm_specific.get("learning_rate") {
if !lr.is_empty() {
self.learning_rate = lr[0];
}
}
if let Some(b1) = state.algorithm_specific.get("beta1") {
if !b1.is_empty() {
self.beta1 = b1[0];
}
}
if let Some(b2) = state.algorithm_specific.get("beta2") {
if !b2.is_empty() {
self.beta2 = b2[0];
}
}
if let Some(eps) = state.algorithm_specific.get("epsilon") {
if !eps.is_empty() {
self.epsilon = eps[0];
}
}
if let Some(wd) = state.algorithm_specific.get("weight_decay") {
if !wd.is_empty() {
self.weight_decay = wd[0];
}
}
if let Some(s) = state.algorithm_specific.get("step") {
if !s.is_empty() {
self.step = s[0].to_usize().unwrap_or(0);
}
}
}
fn set_callback(&mut self, callback: TrainingCallback<T>) {
self.callback = Some(callback);
}
fn call_callback(
&mut self,
epoch: usize,
network: &Network<T>,
data: &TrainingData<T>,
) -> bool {
let error = self.calculate_error(network, data);
if let Some(ref mut callback) = self.callback {
callback(epoch, error)
} else {
true
}
}
}
#[cfg(not(feature = "gpu"))]
pub type GpuAdam<T> = super::Adam<T>;
#[cfg(not(feature = "gpu"))]
pub type GpuAdamW<T> = super::AdamW<T>;
#[cfg(not(feature = "gpu"))]
pub type GpuBatchBackprop<T> = super::BatchBackprop<T>;
#[cfg(not(feature = "gpu"))]
pub type GpuPerformanceStats = ();
pub fn is_gpu_available() -> bool {
#[cfg(miri)]
return false;
#[cfg(feature = "gpu")]
{
crate::webgpu::WebGPUBackend::<f32>::is_available()
}
#[cfg(not(feature = "gpu"))]
false
}
pub fn get_gpu_capabilities() -> String {
#[cfg(miri)]
return "GPU unavailable under Miri".to_string();
#[cfg(feature = "gpu")]
{
if is_gpu_available() {
match crate::webgpu::WebGPUBackend::<f32>::new() {
Ok(backend) => {
let caps = backend.capabilities();
format!(
"GPU Capabilities: Max Buffer {} MB, Compute Units: {}, F16: {}, Bandwidth: {} GB/s",
caps.max_buffer_size / (1024 * 1024),
caps.max_compute_units,
if caps.supports_f16 { "Yes" } else { "No" },
caps.memory_bandwidth_gbps
)
}
Err(_) => "GPU unavailable".to_string(),
}
} else {
"No GPU adapter found".to_string()
}
}
#[cfg(not(feature = "gpu"))]
{
"GPU support not compiled (enable with --features gpu)".to_string()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[cfg_attr(miri, ignore = "Miri cannot handle WebGPU FFI calls")]
fn test_gpu_availability() {
println!("GPU available: {}", is_gpu_available());
println!("GPU capabilities: {}", get_gpu_capabilities());
}
#[cfg(feature = "gpu")]
#[test]
#[cfg_attr(miri, ignore = "Miri cannot handle WebGPU FFI calls")]
fn test_gpu_adam_creation() {
if !is_gpu_available() {
println!("GPU not available, skipping GPU Adam creation test");
return;
}
let result = GpuAdam::new(0.001f32);
match result {
Ok(optimizer) => {
assert_eq!(optimizer.learning_rate, 0.001);
assert_eq!(optimizer.beta1, 0.9);
assert_eq!(optimizer.beta2, 0.999);
assert_eq!(optimizer.step, 0);
}
Err(e) => {
println!("GPU Adam creation failed (expected in CI): {e}");
}
}
}
}