Files
rustytorch/crates/core/rtx-nn/src/utils.rs
T
2026-03-04 00:08:42 +00:00

141 lines
4.0 KiB
Rust

//! 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> {
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> {
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> {
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<f32> {
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<f32> {
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<usize>,
pub dtype: rtx_tensor::DType,
}
/// Calculate tensor statistics
pub fn tensor_stats(tensor: &Tensor) -> Result<TensorStats> {
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());
}