141 lines
4.0 KiB
Rust
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());
|
|
}
|