use crate::OptimizerResult;
use std::ops::Add;
use torsh_core::error::Result;
use torsh_tensor::Tensor;
#[allow(clippy::too_many_arguments)]
pub fn fused_adam_step(
param: &mut Tensor,
grad: &Tensor,
exp_avg: &mut Tensor,
exp_avg_sq: &mut Tensor,
lr: f32,
beta1: f32,
beta2: f32,
eps: f32,
step: u64,
weight_decay: Option<f32>,
) -> Result<()> {
let effective_grad = if let Some(wd) = weight_decay {
let decay_term = param.mul_scalar(wd)?;
grad.add(&decay_term)?
} else {
grad.clone()
};
exp_avg.mul_scalar_(beta1)?;
let grad_term = effective_grad.mul_scalar(1.0 - beta1)?;
*exp_avg = exp_avg.add(&grad_term)?;
exp_avg_sq.mul_scalar_(beta2)?;
let grad_sq = effective_grad.mul_op(&effective_grad)?;
let grad_sq_term = grad_sq.mul_scalar(1.0 - beta2)?;
*exp_avg_sq = exp_avg_sq.add(&grad_sq_term)?;
let bias_correction1 = 1.0 - beta1.powi(step as i32);
let bias_correction2 = 1.0 - beta2.powi(step as i32);
let corrected_lr = lr * (bias_correction2.sqrt()) / bias_correction1;
let denom = exp_avg_sq.sqrt()?.add_scalar(eps)?;
let update = exp_avg.div(&denom)?.mul_scalar(corrected_lr)?;
*param = param.sub(&update)?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fused_sgd_step(
param: &mut Tensor,
grad: &Tensor,
momentum_buffer: Option<&mut Tensor>,
lr: f32,
momentum: f32,
dampening: f32,
weight_decay: Option<f32>,
nesterov: bool,
) -> Result<()> {
let effective_grad = if let Some(wd) = weight_decay {
let decay_term = param.mul_scalar(wd)?;
grad.add(&decay_term)?
} else {
grad.clone()
};
let update = if let Some(buf) = momentum_buffer {
buf.mul_scalar_(momentum)?;
let grad_term = effective_grad.mul_scalar(1.0 - dampening)?;
*buf = buf.add(&grad_term)?;
if nesterov {
let momentum_term = buf.mul_scalar(momentum)?;
momentum_term.add(&effective_grad)?
} else {
buf.clone()
}
} else {
effective_grad
};
let scaled_update = update.mul_scalar(lr)?;
*param = param.sub(&scaled_update)?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fused_rmsprop_step(
param: &mut Tensor,
grad: &Tensor,
square_avg: &mut Tensor,
lr: f32,
alpha: f32,
eps: f32,
weight_decay: Option<f32>,
momentum_buffer: Option<&mut Tensor>,
momentum: f32,
) -> Result<()> {
let effective_grad = if let Some(wd) = weight_decay {
let decay_term = param.mul_scalar(wd)?;
grad.add(&decay_term)?
} else {
grad.clone()
};
square_avg.mul_scalar_(alpha)?;
let grad_sq = effective_grad.mul_op(&effective_grad)?;
let grad_sq_term = grad_sq.mul_scalar(1.0 - alpha)?;
*square_avg = square_avg.add(&grad_sq_term)?;
let denom = square_avg.sqrt()?.add_scalar(eps)?;
let base_update = effective_grad.div(&denom)?;
let update = if let Some(buf) = momentum_buffer {
buf.mul_scalar_(momentum)?;
*buf = buf.add(&base_update)?;
buf.clone()
} else {
base_update
};
let scaled_update = update.mul_scalar(lr)?;
*param = param.sub(&scaled_update)?;
Ok(())
}
pub fn fused_adagrad_step(
param: &mut Tensor,
grad: &Tensor,
sum_of_squares: &mut Tensor,
lr: f32,
eps: f32,
weight_decay: Option<f32>,
) -> Result<()> {
let effective_grad = if let Some(wd) = weight_decay {
let decay_term = param.mul_scalar(wd)?;
grad.add(&decay_term)?
} else {
grad.clone()
};
let grad_sq = effective_grad.mul_op(&effective_grad)?;
*sum_of_squares = sum_of_squares.add(&grad_sq)?;
let denom = sum_of_squares.sqrt()?.add_scalar(eps)?;
let update = effective_grad.div(&denom)?.mul_scalar(lr)?;
*param = param.sub(&update)?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fused_adadelta_step(
param: &mut Tensor,
grad: &Tensor,
square_avg: &mut Tensor,
acc_delta: &mut Tensor,
rho: f32,
eps: f32,
weight_decay: Option<f32>,
) -> Result<()> {
let effective_grad = if let Some(wd) = weight_decay {
let decay_term = param.mul_scalar(wd)?;
grad.add(&decay_term)?
} else {
grad.clone()
};
square_avg.mul_scalar_(rho)?;
let grad_sq = effective_grad.mul_op(&effective_grad)?;
let grad_sq_term = grad_sq.mul_scalar(1.0 - rho)?;
*square_avg = square_avg.add(&grad_sq_term)?;
let rms_grad = square_avg.add_scalar(eps)?.sqrt()?;
let rms_delta = acc_delta.add_scalar(eps)?.sqrt()?;
let delta = effective_grad
.mul_op(&rms_delta)?
.div(&rms_grad)?
.mul_scalar(-1.0)?;
acc_delta.mul_scalar_(rho)?;
let delta_sq = delta.mul_op(&delta)?;
let delta_sq_term = delta_sq.mul_scalar(1.0 - rho)?;
*acc_delta = acc_delta.add(&delta_sq_term)?;
*param = param.add(&delta)?;
Ok(())
}
pub trait FusedKernelSupport {
fn set_fused(&mut self, fused: bool);
fn is_fused(&self) -> bool;
fn fused_stats(&self) -> FusedStats;
}
#[derive(Debug, Clone, Default)]
pub struct FusedStats {
pub fused_ops_count: u64,
pub unfused_ops_count: u64,
pub total_kernel_launches: u64,
pub memory_bandwidth_saved: f64, }
impl FusedStats {
pub fn fusion_efficiency(&self) -> f64 {
if self.fused_ops_count + self.unfused_ops_count == 0 {
0.0
} else {
self.fused_ops_count as f64 / (self.fused_ops_count + self.unfused_ops_count) as f64
* 100.0
}
}
pub fn reset(&mut self) {
*self = Self::default();
}
}
pub mod utils {
use super::*;
pub fn can_fuse_tensors(tensors: &[&Tensor]) -> bool {
if tensors.is_empty() {
return false;
}
let first_device = tensors[0].device();
let first_dtype = tensors[0].dtype();
let first_shape = tensors[0].shape();
tensors.iter().all(|tensor| {
tensor.device() == first_device
&& tensor.dtype() == first_dtype
&& tensor.shape() == first_shape
})
}
pub fn estimate_bandwidth_savings(
tensor_size: usize,
element_size: usize,
num_separate_ops: usize,
num_fused_ops: usize,
) -> f64 {
let separate_bandwidth = tensor_size * element_size * num_separate_ops * 2; let fused_bandwidth = tensor_size * element_size * num_fused_ops * 2; (separate_bandwidth - fused_bandwidth) as f64 / 1e9 }
pub fn supports_fused_ops(device: &dyn torsh_core::device::Device) -> bool {
match device.device_type() {
torsh_core::device::DeviceType::Cpu => true, torsh_core::device::DeviceType::Cuda(_) => true, torsh_core::device::DeviceType::Metal(_) => true, torsh_core::device::DeviceType::Wgpu(_) => false, }
}
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_core::device::Device;
use torsh_tensor::creation;
#[test]
fn test_fused_adam_step() -> OptimizerResult<()> {
let mut param = creation::ones(&[2, 2]).unwrap();
let grad = creation::ones(&[2, 2]).unwrap();
let mut exp_avg = creation::zeros(&[2, 2]).unwrap();
let mut exp_avg_sq = creation::zeros(&[2, 2]).unwrap();
let result = fused_adam_step(
&mut param,
&grad,
&mut exp_avg,
&mut exp_avg_sq,
0.01,
0.9,
0.999,
1e-8,
1,
None,
);
assert!(result.is_ok());
let param_vals = param.to_vec()?;
assert!(param_vals.iter().all(|&x| x < 1.0));
let exp_avg_vals = exp_avg.to_vec()?;
assert!(exp_avg_vals.iter().any(|&x| x != 0.0));
let exp_avg_sq_vals = exp_avg_sq.to_vec()?;
assert!(exp_avg_sq_vals.iter().any(|&x| x != 0.0));
Ok(())
}
#[test]
fn test_fused_sgd_step() -> OptimizerResult<()> {
let mut param = creation::ones(&[2, 2]).unwrap();
let grad = creation::ones(&[2, 2]).unwrap();
let mut momentum_buffer = creation::zeros(&[2, 2]).unwrap();
let result = fused_sgd_step(
&mut param,
&grad,
Some(&mut momentum_buffer),
0.01,
0.9,
0.0,
None,
false,
);
assert!(result.is_ok());
let param_vals = param.to_vec()?;
assert!(param_vals.iter().all(|&x| x < 1.0));
let momentum_vals = momentum_buffer.to_vec()?;
assert!(momentum_vals.iter().any(|&x| x != 0.0));
Ok(())
}
#[test]
fn test_can_fuse_tensors() {
let tensor1 = creation::ones(&[2, 2]).unwrap();
let tensor2 = creation::zeros(&[2, 2]).unwrap();
let tensor3 = creation::ones(&[3, 3]).unwrap();
assert!(utils::can_fuse_tensors(&[&tensor1, &tensor2]));
assert!(!utils::can_fuse_tensors(&[&tensor1, &tensor3]));
assert!(!utils::can_fuse_tensors(&[]));
}
#[test]
fn test_fused_stats() {
let mut stats = FusedStats::default();
stats.fused_ops_count = 80;
stats.unfused_ops_count = 20;
assert_eq!(stats.fusion_efficiency(), 80.0);
stats.reset();
assert_eq!(stats.fused_ops_count, 0);
assert_eq!(stats.unfused_ops_count, 0);
}
#[test]
fn test_estimate_bandwidth_savings() {
let savings = utils::estimate_bandwidth_savings(1000, 4, 5, 1);
assert!(savings > 0.0);
}
#[test]
fn test_supports_fused_ops() {
let cpu_device =
torsh_core::device::DeviceFactory::create_device(torsh_core::device::DeviceType::Cpu)
.unwrap();
assert!(utils::supports_fused_ops(cpu_device.as_ref()));
}
}