//! Utility functions for neural network operations use crate::Result; use rtx_tensor::Tensor; /// Calculate the number of parameters in a tensor pub fn count_parameters(tensor: &Tensor) -> usize { tensor.numel() } /// Calculate total number of parameters in a list of tensors pub fn count_total_parameters(tensors: &[&Tensor]) -> usize { tensors.iter().map(|t| count_parameters(t)).sum() } /// Convert tensor to half precision (FP16) #[cfg(feature = "f16")] pub fn to_half(tensor: &Tensor) -> Result { tensor.to_dtype(rtx_tensor::DType::F16).map_err(Into::into) } /// Convert tensor to single precision (FP32) pub fn to_float(tensor: &Tensor) -> Result { tensor.to_dtype(rtx_tensor::DType::F32).map_err(Into::into) } /// Convert tensor to double precision (FP64) pub fn to_double(tensor: &Tensor) -> Result { tensor.to_dtype(rtx_tensor::DType::F64).map_err(Into::into) } /// Clip gradients by norm pub fn clip_grad_norm(parameters: &mut [&mut Tensor], max_norm: f32) -> Result { let mut total_norm = 0.0_f32; // Calculate total norm for param in parameters.iter() { if let Some(grad) = param.grad() { // Calculate L2 norm: sqrt(sum(x^2)) let squared = grad.mul(&grad)?; let sum = squared.sum(None)?; let param_norm = sum.sqrt()?.item()?; total_norm += param_norm * param_norm; } } total_norm = total_norm.sqrt(); let clip_coef = max_norm / (total_norm + 1e-6); if clip_coef < 1.0 { for param in parameters.iter_mut() { if let Some(grad) = param.grad() { let scaled_grad = grad.mul_scalar(clip_coef)?; param.set_grad(Some(scaled_grad)); } } } Ok(total_norm) } /// Clip gradients by value pub fn clip_grad_value(parameters: &mut [&mut Tensor], clip_value: f32) -> Result<()> { for param in parameters.iter_mut() { if let Some(grad) = param.grad() { let clipped = grad.clamp(-clip_value, clip_value)?; param.set_grad(Some(clipped)); } } Ok(()) } /// Calculate gradient norm pub fn grad_norm(parameters: &[&Tensor]) -> Result { let mut total_norm = 0.0_f32; for param in parameters { if let Some(grad) = param.grad() { // Calculate L2 norm: sqrt(sum(x^2)) let squared = grad.mul(&grad)?; let sum = squared.sum(None)?; let param_norm = sum.sqrt()?.item()?; total_norm += param_norm * param_norm; } } Ok(total_norm.sqrt()) } /// Zero gradients for a list of parameters pub fn zero_grad(parameters: &mut [&mut Tensor]) -> Result<()> { for param in parameters.iter_mut() { param.set_grad(None); } Ok(()) } /// Get model size in bytes pub fn model_size_bytes(parameters: &[&Tensor]) -> usize { parameters.iter().map(|t| t.numel() * 4).sum() // Assuming f32 } /// Get model size in MB pub fn model_size_mb(parameters: &[&Tensor]) -> f64 { model_size_bytes(parameters) as f64 / (1024.0 * 1024.0) } /// Summary statistics for a tensor #[derive(Debug)] pub struct TensorStats { pub mean: f32, pub std: f32, pub min: f32, pub max: f32, pub shape: Vec, pub dtype: rtx_tensor::DType, } /// Calculate tensor statistics pub fn tensor_stats(tensor: &Tensor) -> Result { Ok(TensorStats { mean: tensor.mean(&[], false)?.item()?, std: tensor.var(&[], false, false)?.sqrt()?.item()?, min: tensor.max_keepdim(None, false)?.item()?, max: tensor.max()?.item()?, shape: tensor.shape().dims().to_vec(), dtype: tensor.dtype(), }) } /// Print model summary pub fn print_model_summary(name: &str, parameters: &[&Tensor]) { let total_params = count_total_parameters(parameters); let model_size = model_size_mb(parameters); println!("Model: {name}"); println!("Total parameters: {total_params}"); println!("Model size: {model_size:.2} MB"); println!("Number of layers: {}", parameters.len()); }