use super::{AdamWOptions, FusedAdamWError, StepControl, config::checked_elements, kernel};
use ruda_core::{device::Device, tensor::DType};
use ruda_kernel::{
dsl::{Runtime, calculate_ruda_count_elemwise, prelude::RudaDim},
tensor::{RudaTensor, allocation::empty_device_contiguous_dtype},
};
pub struct AdamWState<R: Runtime> {
step: u64,
first: RudaTensor<R>,
second: RudaTensor<R>,
maximum: Option<RudaTensor<R>>,
}
impl<R: Runtime> Clone for AdamWState<R> {
fn clone(&self) -> Self {
Self { step: self.step, first: self.first.clone(), second: self.second.clone(), maximum: self.maximum.clone() }
}
}
impl<R: Runtime> AdamWState<R> {
pub fn from_parts(
step: u64, first: RudaTensor<R>, second: RudaTensor<R>, maximum: Option<RudaTensor<R>>,
) -> Result<Self, FusedAdamWError> {
if step == 0 { return Err(FusedAdamWError::InvalidState("restored step must be positive")); }
Ok(Self { step, first, second, maximum })
}
pub fn step(&self) -> u64 { self.step }
pub fn first_moment(&self) -> &RudaTensor<R> { &self.first }
pub fn second_moment(&self) -> &RudaTensor<R> { &self.second }
pub fn max_second_moment(&self) -> Option<&RudaTensor<R>> { self.maximum.as_ref() }
pub fn into_parts(self) -> (u64, RudaTensor<R>, RudaTensor<R>, Option<RudaTensor<R>>) {
(self.step, self.first, self.second, self.maximum)
}
}
pub struct AdamWUpdate<R: Runtime> {
pub parameters: RudaTensor<R>,
pub state: Option<AdamWState<R>>,
pub updated: bool,
}
pub fn adamw_step<R: Runtime>(
parameters: &RudaTensor<R>,
gradients: &RudaTensor<R>,
state: Option<&AdamWState<R>>,
options: &AdamWOptions,
control: StepControl,
) -> Result<AdamWUpdate<R>, FusedAdamWError> {
adamw_step_scaled(parameters, gradients, state, options, control, 1.0)
}
pub(super) fn adamw_step_scaled<R: Runtime>(
parameters: &RudaTensor<R>,
gradients: &RudaTensor<R>,
state: Option<&AdamWState<R>>,
options: &AdamWOptions,
control: StepControl,
clip_multiplier: f32,
) -> Result<AdamWUpdate<R>, FusedAdamWError> {
if !clip_multiplier.is_finite() || !(0.0..=1.0).contains(&clip_multiplier) {
return Err(FusedAdamWError::InvalidOption("clip_multiplier"));
}
let size = validate_step(parameters, gradients, state, options, control)?;
if size == 0 || control.skip_update {
return Ok(AdamWUpdate { parameters: parameters.clone(), state: state.cloned(), updated: false });
}
let coefficients = options.prepare_step(state.map_or(0, |s| s.step), control)?;
let client = parameters.client.clone();
let allocate = || empty_device_contiguous_dtype(
client.clone(), parameters.device.clone(), parameters.meta.shape().clone(), DType::F32,
);
let new_parameters = allocate();
let new_first = allocate();
let new_second = allocate();
let new_maximum = options.amsgrad.then(allocate);
let old_first = state.map_or(parameters, |s| &s.first);
let old_second = state.map_or(parameters, |s| &s.second);
let old_maximum = state.and_then(|s| s.maximum.as_ref()).unwrap_or(parameters);
let maximum_output = new_maximum.as_ref().unwrap_or(&new_second);
let dim = RudaDim::new(client.properties(), size);
kernel::adamw::launch::<R>(
&client, calculate_ruda_count_elemwise(&client, size, dim), dim,
parameters.clone().into_array_arg(), gradients.clone().into_array_arg(),
old_first.clone().into_array_arg(), old_second.clone().into_array_arg(), old_maximum.clone().into_array_arg(),
new_parameters.clone().into_array_arg(), new_first.clone().into_array_arg(),
new_second.clone().into_array_arg(), maximum_output.clone().into_array_arg(),
options.learning_rate, options.beta1, options.beta2, options.epsilon,
coefficients.decay_multiplier, coefficients.inverse_bias1, coefficients.inverse_bias2,
coefficients.inverse_gradient_scale, clip_multiplier,
state.is_some(), options.amsgrad, options.maximize, clip_multiplier != 1.0,
concat!(include_str!("kernel.rs"), include_str!("config.rs")).to_owned(),
gradients.dtype.into(),
);
Ok(AdamWUpdate {
parameters: new_parameters,
state: Some(AdamWState { step: coefficients.step, first: new_first, second: new_second, maximum: new_maximum }),
updated: true,
})
}
pub(super) fn validate_step<R: Runtime>(
parameters: &RudaTensor<R>, gradients: &RudaTensor<R>,
state: Option<&AdamWState<R>>, options: &AdamWOptions, control: StepControl,
) -> Result<usize, FusedAdamWError> {
options.validate()?;
control.validate()?;
let size = validate_tensor(parameters, false)?;
validate_related(parameters, gradients, true, "gradient")?;
if let Some(state) = state {
if state.step == 0 || state.maximum.is_some() != options.amsgrad {
return Err(FusedAdamWError::InvalidState("step/AMSGrad mode"));
}
validate_related(parameters, &state.first, false, "first moment")?;
validate_related(parameters, &state.second, false, "second moment")?;
if let Some(maximum) = &state.maximum {
validate_related(parameters, maximum, false, "max second moment")?;
}
}
if size != 0 && !control.skip_update {
state.map_or(0, |s| s.step).checked_add(1).ok_or(FusedAdamWError::StepOverflow)?;
}
Ok(size)
}
fn validate_related<R: Runtime>(
parameters: &RudaTensor<R>, tensor: &RudaTensor<R>, gradient: bool, name: &'static str,
) -> Result<(), FusedAdamWError> {
validate_tensor(tensor, gradient)?;
if tensor.meta.shape() != parameters.meta.shape() {
return Err(FusedAdamWError::ShapeMismatch(name));
}
if tensor.device.to_id() != parameters.device.to_id()
|| !tensor.client.same_execution_queue(¶meters.client)
{
return Err(FusedAdamWError::DifferentExecutionQueue(name));
}
Ok(())
}
pub(super) fn validate_tensor<R: Runtime>(tensor: &RudaTensor<R>, gradient: bool) -> Result<usize, FusedAdamWError> {
let dtype_ok = if gradient {
matches!(tensor.dtype, DType::F32 | DType::F16 | DType::BF16)
} else { tensor.dtype == DType::F32 };
if !dtype_ok || tensor.qparams.is_some() {
return Err(FusedAdamWError::InvalidTensor("parameters/moments require FP32; gradients require F32/F16/BF16"));
}
let size = checked_elements(tensor.meta.shape())?;
if tensor.meta.shape().len() != tensor.meta.strides().len() {
return Err(FusedAdamWError::InvalidTensor("shape/stride rank mismatch"));
}
if size != 0 && !tensor.is_contiguous() {
return Err(FusedAdamWError::InvalidTensor("explicit contiguous inputs required"));
}
let bytes_per_element = tensor.dtype.size() as u64;
let start = tensor.handle.offset_start.unwrap_or(0);
let end = tensor.handle.offset_end.unwrap_or(0);
let usable = tensor.handle.size().checked_sub(start).and_then(|n| n.checked_sub(end))
.ok_or(FusedAdamWError::InvalidTensor("invalid storage offsets"))?;
if start % bytes_per_element != 0 || usable < size as u64 * bytes_per_element {
return Err(FusedAdamWError::InvalidTensor("misaligned or undersized tensor storage"));
}
Ok(size)
}