Files
panic/src/neural_network/optimizers/optimizer_rmsprop.cpp
T
2026-08-06 18:56:36 +02:00

350 lines
11 KiB
C++

/**++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
*
* PANIC
* Portable Algorithms and Numerics In C++
*
* Scientific computing from scratch, with feeling.
*
* Copyright (c) 2026 Michelle Bausager
*
* This file is part of PANIC.
*
* PANIC is free software licensed under the GNU General Public License v3.0 or later.
* You may redistribute and/or modify it under the terms of the GPL.
*
* PANIC is distributed in the hope that it will be useful, but WITHOUT ANY WARRANTY;
* without even the implied warranty of MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.
* See the LICENSE file for the full license text.
*
* SPDX-License-Identifier: GPL-3.0-or-later
*
*++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
*
* Project Name: PANIC
* Module Name: neural_network
* File Name: optimizer_rmsprop.cpp
* Revision: 0.1.0
* Date: 04-08-2026
* Author: Michelle Bausager
*
* Description:
* Defines the optimizer_rmsprop used in neural network
*
*++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++*/
//---------------------------------------------------------------------------------------------------------------------------
// INCLUDE DESCRIPTION
//---------------------------------------------------------------------------------------------------------------------------
#include <neural_network/optimizers/optimizer_rmsprop.hpp>
#include <config/omp.hpp>
#include <math/mul.hpp>
#include <math/add.hpp>
#include <math/sub.hpp>
#include <math/div.hpp>
#include <math/sqrt.hpp>
//---------------------------------------------------------------------------------------------------------------------------
// PRIVATE CONSTANTS
//---------------------------------------------------------------------------------------------------------------------------
/**
* @brief Minimum number of element operations before using the OpenMP-enabled loop.
*
* Small vectors and matrices are kept serial because the overhead of starting
* worker threads can be larger than the work itself.
*/
static const panic::types::uint_t optimizer_rmsprop_omp_min_size = 250;
static const panic::types::real_t real_t_1 = static_cast<panic::types::real_t>(1);
//---------------------------------------------------------------------------------------------------------------------------
// INPLEMENTATION
//---------------------------------------------------------------------------------------------------------------------------
namespace panic{
namespace neural_network{
//--------------------------------------------------------------------------------------------------------------------------
// Constructor Name : panic::neural_network::optimizer_rmsprop
//
// Description:
// Constructor for optimizer_rmsprop.
//--------------------------------------------------------------------------------------------------------------------------
optimizer_rmsprop::optimizer_rmsprop(const panic::types::real_t learning_rate,
const panic::types::real_t decay,
const panic::types::real_t epsilon,
const panic::types::real_t rho){
this->learning_rate = learning_rate;
this->decay = decay;
this->epsilon = epsilon;
this->rho = rho;
iterations = static_cast<panic::types::uint_t>(0);
current_learning_rate = learning_rate;
}
//--------------------------------------------------------------------------------------------------------------------------
// Constructor Name : panic::neural_network::update_params
//
// Description:
// Updates weights and biases in layer.
//--------------------------------------------------------------------------------------------------------------------------
bool optimizer_rmsprop::update_params(trainable_layer& layer) {
// Gradients must match their corresponding parameters.
if (layer.weights.rows() != layer.dweights.rows() || layer.weights.cols() != layer.dweights.cols() || layer.biases.size() != layer.dbiases.size()){
return false;
}
if (layer.weights.rows() == static_cast<panic::types::uint_t>(0) || layer.weights.cols() == static_cast<panic::types::uint_t>(0) || layer.biases.size() == static_cast<panic::types::uint_t>(0) ){
return false;
}
// If layer doesn't contain cache arrays, create them
// and fill them with zeros
if(layer.weights.rows() != layer.weight_cache.rows() ||
layer.weights.cols() != layer.weight_cache.cols() ||
layer.biases.size() != layer.bias_cache.size()){
if (! layer.weight_cache.resize(layer.weights.rows(), layer.weights.cols())){
return false;
}
layer.weight_cache.fill(0);
if (! layer.bias_cache.resize(layer.biases.size())){
return false;
}
layer.bias_cache.fill(0);
}
//--------------------------------------------------------------------------------------------------------------------------
// Update caches with squared current gradients
//--------------------------------------------------------------------------------------------------------------------------
const panic::types::real_t one_minus_rho =
static_cast<panic::types::real_t>(1) - rho;
//--------------------------------------------------------------------------------------------------------------------------
// Weight cache
//--------------------------------------------------------------------------------------------------------------------------
panic::tensor::real_matrix retained_weight_cache;
panic::tensor::real_matrix squared_dweights;
panic::tensor::real_matrix weighted_dweights;
if (!panic::math::mul(
layer.weight_cache,
rho,
retained_weight_cache
)){
return false;
}
if (!panic::math::mul(
layer.dweights,
layer.dweights,
squared_dweights
)){
return false;
}
if (!panic::math::mul(
squared_dweights,
one_minus_rho,
weighted_dweights
)){
return false;
}
if (!panic::math::add(
retained_weight_cache,
weighted_dweights,
layer.weight_cache
)){
return false;
}
//--------------------------------------------------------------------------------------------------------------------------
// Bias cache
//--------------------------------------------------------------------------------------------------------------------------
panic::tensor::real_vector retained_bias_cache;
panic::tensor::real_vector squared_dbiases;
panic::tensor::real_vector weighted_dbiases;
if (!panic::math::mul(
layer.bias_cache,
rho,
retained_bias_cache
)){
return false;
}
if (!panic::math::mul(
layer.dbiases,
layer.dbiases,
squared_dbiases
)){
return false;
}
if (!panic::math::mul(
squared_dbiases,
one_minus_rho,
weighted_dbiases
)){
return false;
}
if (!panic::math::add(
retained_bias_cache,
weighted_dbiases,
layer.bias_cache
)){
return false;
}
//--------------------------------------------------------------------------------------------------------------------------
// Calculate weight updates
//--------------------------------------------------------------------------------------------------------------------------
panic::tensor::real_matrix weight_updates;
if (!panic::math::mul(
layer.dweights,
-current_learning_rate,
weight_updates
)){
return false;
}
panic::tensor::real_matrix sqrt_weight_cache;
panic::tensor::real_matrix weight_denominator;
if (!panic::math::sqrt(
layer.weight_cache,
sqrt_weight_cache
)){
return false;
}
if (!panic::math::add(
sqrt_weight_cache,
epsilon,
weight_denominator
)){
return false;
}
if (!panic::math::div(
weight_updates,
weight_denominator,
weight_updates
)){
return false;
}
if (!panic::math::add(
layer.weights,
weight_updates,
layer.weights
)){
return false;
}
//--------------------------------------------------------------------------------------------------------------------------
// Calculate bias updates
//--------------------------------------------------------------------------------------------------------------------------
panic::tensor::real_vector bias_updates;
if (!panic::math::mul(
layer.dbiases,
-current_learning_rate,
bias_updates
)){
return false;
}
panic::tensor::real_vector sqrt_bias_cache;
panic::tensor::real_vector bias_denominator;
if (!panic::math::sqrt(
layer.bias_cache,
sqrt_bias_cache
)){
return false;
}
if (!panic::math::add(
sqrt_bias_cache,
epsilon,
bias_denominator
)){
return false;
}
if (!panic::math::div(
bias_updates,
bias_denominator,
bias_updates
)){
return false;
}
if (!panic::math::add(
layer.biases,
bias_updates,
layer.biases
)){
return false;
}
return true;
}
//--------------------------------------------------------------------------------------------------------------------------
// Constructor Name : panic::neural_network::pre_update_params
//
// Description:
// function to update internal parameters before update_params()
//--------------------------------------------------------------------------------------------------------------------------
bool optimizer_rmsprop::pre_update_params(){
if (decay){
current_learning_rate = learning_rate * (real_t_1 / (real_t_1 + (decay * iterations)));
}
return true;
}
//--------------------------------------------------------------------------------------------------------------------------
// Constructor Name : panic::neural_network::post_update_params
//
// Description:
// function to update internal parameters after update_params()
//--------------------------------------------------------------------------------------------------------------------------
bool optimizer_rmsprop::post_update_params(){
iterations += static_cast<panic::types::uint_t>(1);
return true;
}
} // namespace tensor
} // namespace panic